
    xi                       d Z ddlZddlZddlZddlZddlZddlZddlmZm	Z	m
Z
mZmZmZ ddlmZ ddlZddlZddlmZ ddlmZ ddlmZmZ ddlmZ dd	lmZmZ dd
lmZ ddlm Z m!Z!m"Z"m#Z#m$Z$m%Z% ddl&m'Z'm(Z(m)Z)m*Z* dZ+d Z, G d d      Z-y)zk
5Hz LM (Language Model) Handler
Handles all LM-related operations including initialization and generation
    N)OptionalDictAnyTupleListUnion)contextmanager)logger)tqdm)AutoTokenizerAutoModelForCausalLM)BaseStreamer)LogitsProcessorList RepetitionPenaltyLogitsProcessor)"MetadataConstrainedLogitsProcessor)DEFAULT_LM_INSTRUCTION!DEFAULT_LM_UNDERSTAND_INSTRUCTIONDEFAULT_LM_INSPIRED_INSTRUCTIONDEFAULT_LM_REWRITE_INSTRUCTIONDURATION_MINDURATION_MAX)get_lm_gpu_memory_ratioget_gpu_memory_gbget_lm_model_sizeget_global_gpu_configg       @c            
          t         j                  } t        | dd      dk7  rnt         j                  j	                  d      rNt        j                  dt         j                  j                         d    dt        | dd       dt        d	
       y y y )NreleaselevelfinallinuxzDetected pre-release Python r   z ( z). This is known to cause segmentation faults with vLLM/nano-vllm on Linux. Please install a stable Python release (e.g. 3.11.12+), or use --backend pt as a workaround.   )
stacklevel)
sysversion_infogetattrplatform
startswithwarningswarnversionsplitRuntimeWarning)vs    4/mnt/workspace/ACE-Step-1.5/acestep/llm_inference.py_warn_if_prerelease_pythonr/   "   s    Aq.'*g5#,,:Q:QRY:Z*3;;+<+<+>q+A*B"WQP^`bEcDd ek k 	
 ;[5    c            2          e Zd ZdZdZej                  j                  d      duZdyde	e
   fdZdzdZdzd	Zde
fd
Zdee
   fdZd{de
dedededeeef   f
dZ	 dyde	e   de
de	e   defdZde
defdZdedefdZ	 	 	 d|dedede	e   de	ee
e	e
   f      dedededede
d ed!e	e   d"e	e   de	e   fd#Z	 d}d$e
d%e
d&e
de
de
d ede
fd'Zde
d(e
deee
f   fd)Zd*ej@                  d+e	e   dej@                  fd,Z!d*ej@                  d-e	e   dej@                  fd.Z"d*ej@                  d/edej@                  fd0Z#d1ej@                  d2ed3e	e   defd4Z$d5e	e   d1ej@                  fd6Z%d7e&d8ej@                  d9ee
e&f   d:e	e&   d;ede&fd<Z'd=e(e
ee
   f   deee
   ef   fd>Z)	 	 	 	 d~d?e
d@e
dAe
d(e
dBedCe	ejT                     dee
ef   fdDZ+d}de
dEede
fdFZ,	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd=e(e
ee
   f   d/edGede
d+e	e   d-e	e   dededed!e	e   d"e	e   de	e   de	ee
e	e
   f      dedededede
d$e
d%e
d&e
dHe	ee      de(e
ee
   f   f.dIZ-dJe
d/edGede
d+e	e   d-e	e   dededede	e   de	ee
e	e
   f      dedededede
d$e
d%e
d&e
de
f(dKZ.	 	 	 	 	 	 	 	 	 	 	 	 	 dd=e(e
ee
   f   d/edGede
d+e	e   d-e	e   dededede	e   de	ee
e	e
   f      dedededede
d$e
d%e
d&e
dHe	ee      de(e
ee
   f   f*dLZ/de	ee
e	e
   f      defdMZ0dNee
e&f   de
fdOZ1	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd$e
d%e
dPe
d/edGede
d+e	e   d-e	e   dededede	e   de	ee
e	e
   f      dQedRedSedTe	e   dHe	ee      dee
e&f   f&dUZ2dd$e
d%e
dVede
de
de
fdWZ3dd$e
d%e
d&e
dVede
de
fdXZ4	 	 ddYe
dVede
de
fdZZ5	 	 	 	 	 	 ddYe
d/ed+e	e   d-e	e   dedededeee
e&f   e
f   fd[Z6d\e
de
fd]Z7	 	 	 dd^e
d_edVede
de
f
d`Z8	 	 	 	 	 	 	 	 dd^e
d_edae	e
   d/ed+e	e   d-e	e   dedededeee
e&f   e
f   fdbZ9	 	 dd$e
d%e
dVede
de
f
dcZ:	 	 	 	 	 	 	 dd$e
d%e
de	ee
e&f      d/ed+e	e   d-e	e   dedededeee
e&f   e
f   fddZ;	 	 	 	 ddJe
dee	ee
e&f      dedededee
e
f   fdfZ<	 dydgej@                  dhe	ej@                     died/ed+e	e   d-e	e   ded3edje	e=   d5e	e   dej@                  fdkZ>	 dydlej@                  dme	ej@                     died/edGed+e	e   d-e	e   ded3edje	e=   d5e	e   dej@                  fdnZ?d\e
deee
e&f   e
f   fdoZ@eAdefdp       ZBde
deee
f   fdqZCdr ZD	 dydJe
dTed/edGede
d+e	e   d-e	e   dededede	e   d$e
d%e
d&e
dHe	ee      dee
   f dsZEdJe
d/edGede
d+e	e   d-e	e   dededede	e   de	ee
e	e
   f      dedededede
d$e
d%e
d&e
de
f(dtZFdJe
d/edGede
d+e	e   d-e	e   dededede	e   de	ee
e	e
   f      dedededede
d$e
d%e
d&e
de
f(duZG	 	 	 	 	 	 	 	 	 	 	 	 	 dd=e(e
ee
   f   d/edGede
d+e	e   d-e	e   dededede	e   de	ee
e	e
   f      dedededede
d$e
d%e
d&e
dHe	ee      de(e
ee
   f   f*dvZHeIdw        ZJdx ZKy)
LLMHandlerz(5Hz LM Handler for audio code generation</think>SPACE_IDNpersistent_storage_pathc                    d| _         d| _        d| _        d| _        d| _        d| _        t        j                  | _        d| _	        t        j                  j                  dd      j                         dv xs; t        t        j                   d      xr t        j                   j#                          | _        || j&                  rd	}|| _        d| _        d| _        d| _        d| _        y)
z)Initialize LLMHandler with default valuesNF   cpuACESTEP_DISABLE_TQDMr    )1trueyesisattyz/data)llmllm_tokenizerllm_initializedllm_backendmax_model_lendevicetorchfloat32dtypeoffload_to_cpuosenvirongetlowerhasattrr#   stderrr=   disable_tqdmIS_HUGGINGFACE_SPACEr5   constrained_processor_hf_model_for_scoring
_mlx_model_mlx_model_path)selfr5   s     r.   __init__zLLMHandler.__init__6   s   !$!]]
#JJNN+A2FLLNRff  fovwz  xB  xB  DL  pM  pe  RU  R\  R\  Rc  Rc  Re  kf #*t/H/H&-#'>$ TX" &*" #r0   returnc                 r   	 | j                   dk(  rA	 t        | j                  d      r| j                  j                          | j                          d| _        d| _        d| _        d| _        d| _         d| _	        d| _
        	 ddl}|j                          t        j                  j                         r=t        j                  j!                          t        j                  j#                          yt        t        d      rt        j$                  j&                  j                         rqt        t        j&                  d      rt        j&                  j#                          t        t        j&                  d      rt        j&                  j!                          yt        t        d	      r\t        j(                  j                         r=t        j(                  j!                          t        j(                  j#                          yyyy# t        $ r Y w xY w# t        $ r Y w xY w# t        $ r Y yw xY w)
z=Release LM weights/tokenizer and clear caches to free memory.vllmresetNFr   mpssynchronizeempty_cachexpu)rA   rL   r>   rY   	Exception _cleanup_torch_distributed_stater?   rP   r@   rR   rS   gccollectrD   cudais_availabler\   r[   backendsrZ   r]   )rT   r`   s     r.   unloadzLLMHandler.unloadQ   s    	6)txx1( 557DH!%D)-D&#(D #D"DO#'D 

 zz&&(

&&(

&&(&5>>+=+=+J+J+L599m4II))+599m4II))+&599+A+A+C		%%'		%%' ,D& 5+ !     		sh   H* 0H
 AH* H AH* 3B(H* A*H* 
	HH* HH* 	H'#H* &H''H* *	H65H6c                     	 ddl m} |j                         r7|j                         r&t	        j
                  d       |j                          yyy# t        $ r"}t	        j
                  d|        Y d}~yd}~ww xY w)zIDestroy default torch distributed process group when already initialized.r   NzM[LLM vLLM] Destroying stale default process group before/after vLLM lifecyclez4[LLM vLLM] Failed to clean torch distributed state: )torch.distributeddistributedrc   is_initializedr
   warningdestroy_process_groupr^   )rT   distexcs      r.   r_   z+LLMHandler._cleanup_torch_distributed_stateu   sp    	Y,  "t':':'<no**, (="  	YNNQRUQVWXX	Ys   AA 	A;A66A;c                 l   | j                   r*t        j                  j                  | j                   d      S t        j                  j	                  t
              }t        j                  j                  t        j                  j                  |            }t        j                  j                  |d      S )z9Get checkpoint directory, prioritizing persistent storagecheckpoints)r5   rH   pathjoinabspath__file__dirname)rT   current_fileproject_roots      r.   _get_checkpoint_dirzLLMHandler._get_checkpoint_dir   sl    ''77<< < <mLLwwx0wwrww|'DEww||L-88r0   c                    | j                         }g }t        j                  j                  |      r}t        j                  |      D ]e  }t        j                  j                  ||      }t        j                  j                  |      sC|j                  d      sU|j                  |       g |j                          |S )zIScan and return all model directory names starting with 'acestep-5Hz-lm-'zacestep-5Hz-lm-)
rw   rH   rp   existslistdirrq   isdirr'   appendsort)rT   checkpoint_dirmodelsitem	item_paths        r.   get_available_5hz_lm_modelsz&LLMHandler.get_available_5hz_lm_models   s    11377>>.)

>2GGLL>	77==+@Q0RMM$' 3
 	r0   
model_pathminimal_gpu	min_ratio	max_ratioc                 R   	 t        j                  d      }t         j                  j                  |      j                  }|dz  }d}|r>t        ||      \  }	}
t        j                  d| d|
 d|	dd|d	d
	       |dk  rd}|	|fS t         j                  j                  |      }|dz  }||z
  }||k  rd|z  }d}||k\  rt        |t        |||z              }	nt        |t        ||dz  |z              }	|	|fS # t        $ r"}t        j                  d|        Y d}~yd}~ww xY w)a  
        Get GPU memory utilization ratio based on LM model size and available GPU memory.

        Args:
            model_path: LM model path (e.g., "acestep-5Hz-lm-0.6B"). Used to determine target memory.
            minimal_gpu: Minimum GPU memory requirement in GB (fallback)
            min_ratio: Minimum memory utilization ratio
            max_ratio: Maximum memory utilization ratio

        Returns:
            Tuple of (gpu_memory_utilization_ratio, low_gpu_memory_mode)
        zcuda:0   @Fz%Adaptive LM memory allocation: model=z	, target=z
GB, ratio=.3fz, total_gpu=.1fGB   T      ?g?z,Failed to calculate GPU memory utilization: N)?F)rD   rC   rb   get_device_propertiestotal_memoryr   r
   infomemory_reservedminmaxr^   rj   )rT   r   r   r   r   rC   total_gpu_mem_bytes	total_gpulow_gpu_memory_moderatiotarget_memory_gbreserved_mem_bytesreserved_gpuavailable_gpues                  r.   get_gpu_memory_utilizationz%LLMHandler.get_gpu_memory_utilization   su   #	\\(+F"'**"B"B6"J"W"W+g5I"' *A*i*X''CJ<yYiXjjtuz{~t  @L  MV  WZ  L[  []  ^  _ q=*.'111 "'!;!;F!C-7L%4M;&!Io&*#+Is9kI6M'NOIs9}s7Ji6W'XY--- 	NNI!MN	s   BC; A3C; ;	D&D!!D&target_durationgeneration_phasefallback_maxc                 n   |i|dkD  rdt         }	 t               }t        |j                  t               }t        t        t        ||            }t        |dz        }|dk(  r|dz   }n|dz   }n||}nt        | dd      dz
  }t        | d      rt        || j                  dz
        }|S # t        $ r Y w xY w)	a"  
        Compute max_new_tokens based on target duration and generation phase.

        In the two-phase architecture:
        - CoT phase: generates metadata (~50-200 tokens) + needs buffer for safety.
        - Codes phase: CoT is already in the prompt; only audio codes are generated.
          The constrained decoder forces EOS at exactly target_codes, so only a
          small buffer (10 tokens) is needed to avoid a misleading progress bar.

        Duration is clamped to ``[DURATION_MIN, max_dur]`` where *max_dur* is the
        GPU-config-dependent maximum (from ``get_global_gpu_config()``) capped at
        ``DURATION_MAX``.  This keeps the progress-bar total aligned with what the
        constrained decoder actually enforces.

        Args:
            target_duration: Target duration in seconds (5 codes = 1 second).
            generation_phase: "cot" or "codes".
            fallback_max: Fallback value when target_duration is not set.

        Returns:
            Computed max_new_tokens value, capped at model's max length.
        r      codes
   i  rB   r7   @   )r   r   r   max_duration_with_lmr^   r   r   intr%   rL   rB   )	rT   r   r   r   gpu_max_durgpu_cfgeffective_durationtarget_codesmax_new_tokenss	            r.   _compute_max_new_tokensz"LLMHandler._compute_max_new_tokens   s    8 &?Q+> 'K/1!'">">M "%\3{O3T!U1A56L7* ".!2 ".!3'!-!(!E!J 4) 1C1Cb1HIN-  s   $B( (	B43B4negative_promptc                 T    |xr% |j                         xr |j                         dk7  S )z:Check if negative prompt is meaningful (not default/empty)NO USER INPUT)strip)rT   r   s     r.   _has_meaningful_negative_promptz*LLMHandler._has_meaningful_negative_prompt  s*    i?#8#8#:i?T?T?VZi?iir0   repetition_penaltyc                 Z    t               }|dk7  r|j                  t        |             |S )z=Build logits processor list with repetition penalty if needed      ?penalty)r   r|   r   )rT   r   logits_processors      r.   _build_logits_processorz"LLMHandler._build_logits_processor  s.    .0$##$DM_$`ar0   use_constrained_decodingconstrained_decoding_debuguser_metadatastop_at_reasoningskip_genresskip_captionskip_languageis_batchmetadata_temperaturecodes_temperaturec                    |
 xr
 |duxs |du}|s|sy| j                   j                          || j                   _        || j                   _        |r#|| j                   _        || j                   _        n"d| j                   _        d| j                   _        | j                   j                  |       |
r| j                   j                  d       | j                   j                  d       | j                   j                  d       | j                   j                  d       | j                   j                  d       n| j                   j                  |       | j                   j                  |       | j                   j                  |       | j                   j                  |       | j                   j                  |       | j                   j                  |	       | j                   S )z8Setup and configure constrained processor for generationNFT)rP   rY   enableddebugr   r   set_target_durationset_user_metadataset_stop_at_reasoningset_skip_genresset_skip_captionset_skip_languageset_generation_phase)rT   r   r   r   r   r   r   r   r   r   r   r   r   use_phase_temperaturess                 r.   _setup_constrained_processorz'LLMHandler._setup_constrained_processor  s     &.!u3Gt3S3tWhptWt'0F 	""((* .F""*+E""( ">RD&&;;LD&&8>BD&&;;?D&&8""66G &&88>&&<<UC&&66t<&&77=&&88> &&88G&&<<=NO&&66{C&&77E&&88G 	""778HI)))r0   captionlyricscot_textc                 h    |s|dk(  r| j                  |||d|      S | j                  ||dd|      S )zKBuild unconditional prompt for CFG based on generation phase and batch moder   T)is_negative_promptr   cot)r   r   r   )build_formatted_prompt_with_cotbuild_formatted_prompt)rT   r   r   r   r   r   r   s          r.   _build_unconditional_promptz&LLMHandler._build_unconditional_promptH  sY     '7277dTc 8   ..D5bq /  r0   rC   c                 H   	 t        j                  |d      | _        | j                  s:| j                  j	                  |      j	                  | j
                        | _        n9| j                  j	                  d      j	                  | j
                        | _        | j                  j                          d| _        d| _        t        j                  d|        d| d| }d|fS # t        $ r/}dd	t        |       d
t        j                          fcY d}~S d}~ww xY w)zALoad PyTorch model from path and return (success, status_message)T)trust_remote_coder8   ptz95Hz LM initialized successfully using PyTorch backend on +   ✅ 5Hz LM initialized successfully
Model: z
Backend: PyTorch
Device: F   ❌ Error initializing 5Hz LM: 

Traceback:
N)r   from_pretrainedr>   rG   torF   evalrA   r@   r
   r   r^   str	traceback
format_exc)rT   r   rC   
status_msgr   s        r.   _load_pytorch_modelzLLMHandler._load_pytorch_model^  s    	m+;;JZ^_DH&&88;;v.11$**=88;;u-00<HHMMO#D#'D KKSTZS[\]G
|SopvowxJ## 	m;CF8CST]ThThTjSklll	ms   C&C) )	D!2$DD!D!logitstop_kc                 n    |2|dkD  r-|t        j                  ||      d   d   k  }t        d      ||<   |S )zApply top-k filtering to logitsr   ).N-inf)rD   topkfloat)rT   r   r   indices_to_removes       r.   _apply_top_k_filterzLLMHandler._apply_top_k_filtero  sB     &FE)B1)Em)T T(-fF$%r0   top_pc                 b   |d|cxk  rdk  rn |S t        j                  |d      \  }}t        j                  t        j                  |j	                         d      d      }||kD  }|dddf   j                         |dd	df<   d
|d<   |j                  d	||      }t	        d      ||<   |S )z)Apply top-p (nucleus) filtering to logitsN        r   T)
descendingr   dim.   r   ).r   r   )rD   r}   cumsumsoftmaxr   clonescatter)rT   r   r   sorted_logitssorted_indicescumulative_probssorted_indices_to_remover   s           r.   _apply_top_p_filterzLLMHandler._apply_top_p_filterv  s    u!2s!2  -2JJv$,O)M>$||EMM-:M:M:OUW,X^`a'7%'?$0Hcrc0R0X0X0Z$S!"W-/0$V, 8 @ @NTl m(-fF$%r0   temperaturec                     |dkD  rP|j                         |z  }t        j                  |d      }t        j                  |d      j	                  d      S t        j
                  |d      S )zSample tokens from logits with temperature.

        Upcasts to float32 for numerical stability (float16 logits can overflow
        during softmax, especially after CFG scaling).
        r   r   r   r   )num_samples)r   rD   r   multinomialsqueezeargmax)rT   r   r  probss       r.   _sample_tokenszLLMHandler._sample_tokens  sZ     ?\\^k1FMM&b1E$$U:BB1EE<<B//r0   tokenseos_token_idpad_token_idc                 v    t        j                  ||k(        ry|||k7  rt        j                  ||k(        ryy)z3Check if any token in the batch is EOS or pad tokenTF)rD   any)rT   r	  r
  r  s       r.   _check_eos_tokenzLLMHandler._check_eos_token  s:    99V|+,#(Dyy</0r0   rP   c                     |@t        |j                  d         D ]$  }|j                  ||   j                                & yy)z8Update constrained processor state with generated tokensNr   )rangeshapeupdate_stater   )rT   rP   r	  bs       r.   #_update_constrained_processor_statez.LLMHandler._update_constrained_processor_state  s=     ,6<<?+%226!9>>3CD , -r0   modelgenerated_idsmodel_kwargspast_key_values	use_cachec                 Z    | |dd|i|d|i}|S  |d|ddddf   |d|d|i}|S )z*Perform forward pass with KV cache supportN	input_idsr  r   )r  r   r  )rT   r  r  r  r  r  outputss          r.   _forward_passzLLMHandler._forward_pass  sr     " ' $G   '23/ /  $	G r0   formatted_promptsc                 8    t        |t              }|r||fS |g|fS )zPNormalize batch input: convert single string to list and return (list, is_batch))
isinstancelist)rT   r  r   s      r.   _normalize_batch_inputz!LLMHandler._normalize_batch_input  s+    /6$h..%&00r0   r~   lm_model_pathbackendrG   rF   c           	         	 |dk(  rt         j                  j                         rd}nt        t         j                  d      r,t         j                  j
                  j                         rd}nt        t         d      r"t         j                  j                         rd}nd}n|dk(  rt         j                  j                         st        t         j                  d      rAt         j                  j
                  j                         rt        j                  d       d}nt        t         d      r7t         j                  j                         rt        j                  d       d}nt        j                  d       d}n|dk(  rt        t         j                  d      r(t         j                  j
                  j                         st         j                  j                         rt        j                  d	       d}n8t        t         d      r6t         j                  j                         rt        j                  d
       d}nt        j                  d       d}n|dk(  rt        t         d      rt         j                  j                         st         j                  j                         rt        j                  d       d}nqt        t         j                  d      r@t         j                  j
                  j                         rt        j                  d       d}nt        j                  d       d}|| _	        || _
        |0|dv rt         j                  | _        nwt         j                  | _        na|| _        |dk(  rU| j                  t         j                  k7  r8t        j                  d| j                   d       t         j                  | _        |d}t        j                  d|        t        j                   j#                  ||      }t        j                   j%                  |      sd| dfS |dk(  rZt         j                  j                         r<t         j                  j'                          t         j                  j)                          t        j                  d       t+        j*                         }t-        j.                  |d      }	t        j                  dt+        j*                         |z
  dd       |	| _        t        j                  d       t+        j*                         }
t3               }|j4                  }t        j                  d| d|j6                   d        t9        | j0                  dd|!      | _        t        j                  d"t+        j*                         |
z
  dd       t        t         j<                  d#      xr t         j<                  j>                  du}tA        |      }|d$k(  s
|d%k(  r|dk(  r| jC                         rt        j                  d&       | jE                  |      \  }}|r|dfS t        j                  d'|        |d$k(  ryt        j                  d(       | jG                  ||      \  }}|s|dfS d)| d*}|dfS |d$k(  r:t        j                  d+       | jG                  ||      \  }}|s|dfS d,| d*}|dfS |d%k(  r |dk7  rt        j                  d-| d.       d/}|d%k(  rtI                |dk(  r
tK               nd0}d0}|dk(  rt         j                  j                         r	 t        t         j                  d1      r't         j                  jM                         \  }}|d2z  }nNt         j                  jO                  d3      jP                  }|t         j                  jS                  d3      z
  d2z  }|dk(  rP|tV        k  rGt        j                  d4|dd5|dd6tV         d7       | jG                  ||      \  }}|s|dfS d8| d*}n| jY                  ||9      }t        j                  d:|        |j[                  d;      r| j\                  s|dk(  rX| jC                         rHt        j                  d<       | jE                  |      \  }}|r|dfS t        j                  d=| d>       t        j                  d?       | jG                  ||      \  }}|s|dfS d8| d*}n |d$k7  r| jG                  ||      \  }}|s|dfS dfS # tT        $ r d0}Y cw xY w# tT        $ r/}d@t_        |       dAta        jb                          dfcY d}~S d}~ww xY w)Ba  
        Initialize 5Hz LM model

        Args:
            checkpoint_dir: Checkpoint directory path
            lm_model_path: LM model path (relative to checkpoint_dir)
            backend: Backend type ("vllm" or "pt")
            device: Device type ("auto", "cuda", "mps", "xpu", or "cpu")
            offload_to_cpu: Whether to offload to CPU
            dtype: Data type (if None, auto-detect based on device)

        Returns:
            (status_message, success)
        autorb   rZ   r]   r8   zA[initialize] CUDA requested but unavailable. Falling back to MPS.zA[initialize] CUDA requested but unavailable. Falling back to XPU.zA[initialize] CUDA requested but unavailable. Falling back to CPU.zA[initialize] MPS requested but unavailable. Falling back to CUDA.z@[initialize] MPS requested but unavailable. Falling back to XPU.z@[initialize] MPS requested but unavailable. Falling back to CPU.zA[initialize] XPU requested but unavailable. Falling back to CUDA.z@[initialize] XPU requested but unavailable. Falling back to MPS.z@[initialize] XPU requested but unavailable. Falling back to CPU.N)rb   r]   z([initialize] Overriding requested dtype z to float32 for LM on MPS.zacestep-5Hz-lm-1.7Bz3[initialize] lm_model_path is None, using default: u   ❌ 5Hz LM model not found at Fz.loading 5Hz LM tokenizer... it may take 80~90sT)use_fastz(5Hz LM tokenizer loaded successfully in .2f secondsz.Initializing constrained decoding processor...z-Setting constrained decoding max_duration to zs based on GPU config (tier: ))	tokenizerr   r   max_durationz%Constrained processor initialized in hipmlxrX   z8Attempting MLX backend for Apple Silicon acceleration...zMLX backend failed: zDMLX explicitly requested but failed, falling back to PyTorch backendu:   ✅ 5Hz LM initialized (PyTorch fallback from MLX)
Model: z
Backend: PyTorchz;MLX not available (requires Apple Silicon + mlx-lm package)uD   ✅ 5Hz LM initialized (PyTorch fallback, MLX not available)
Model: zJ[initialize] vllm backend requires CUDA, using PyTorch backend for device=.r   r   mem_get_infor   r   z3vLLM disabled due to insufficient free VRAM (total=z	GB, free=z
GB, need>=u,   GB free) — falling back to PyTorch backendu>   ✅ 5Hz LM initialized successfully (PyTorch fallback)
Model: )enforce_eagerz5Hz LM status message: u   ❌z)vllm failed on MPS, trying MLX backend...zMLX also failed: z, falling back to PyTorchzFalling back to PyTorch backendr   r   )2rD   rb   rc   rL   rd   rZ   r]   r
   rj   rC   rG   bfloat16rF   rE   r   rH   rp   rq   ry   r\   r[   timer   r   r?   r   r   tierr   rP   r*   r.  bool_is_mlx_available_load_mlx_modelr   r/   r   r1  r   r   r   r^   VRAM_SAFE_FREE_GB_initialize_5hz_lm_vllmr'   r@   r   r   r   )rT   r~   r$  r%  rC   rG   rF   full_lm_model_path
start_timer?   processor_start
gpu_configmax_duration_for_constraintis_rocmenforce_eager_for_vllmmlx_success
mlx_statussuccessr   total_gbfree_gb
free_bytes_total_bytesr   s                            r.   
initializezLLMHandler.initialize  sr   .~	m::**,#FU^^U38J8J8W8W8Y"FUE*uyy/E/E/G"F"F6!%***A*A*C5>>51enn6H6H6U6U6WNN#fg"FUE*uyy/E/E/GNN#fg"FNN#fg"F5'%..%*HU^^M_M_MlMlMn::**,NN#fg#FUE*uyy/E/E/GNN#ef"FNN#ef"F5'%*?EIIDZDZD\::**,NN#fg#FU^^U38J8J8W8W8YNN#ef"FNN#ef"F DK"0D }_,!&DJ!&DJ"
U?tzzU]]'BNNB4::,Nhi "'DJ $ 5QR_Q`ab!#nm!L77>>"4578J7KLeSS EJJ$;$;$=

&&(

&&(KKHIJ)99:LW[\MKKB499;Q[C[\_B``hij!.D KKHI"iikO.0J*4*I*I'KKGHcGd  eB  CM  CR  CR  BS  ST  U  V)K,,8	*D& KK?		o@]^a?bbjkl emmU3U8I8IQU8UG%)']" %Gv$5&E/))+KK Z[.2.B.BCU.V+K")4//)=j\'JK"e+"NN+qr262J2JK]_e2f/GZ#*'15'8 8+fgyfz  {M  *NJ#-t#33%NN#`a*.*B*BCUW]*^'GZ")500#hi{h|  }O  "PJ%t++& Vv%5`ag`hhij  & *,28F2B,.V#

(?(?(A&"5::~>,1JJ,C,C,EMJ&0G&<G*/***J*J1*M*Z*ZK'2UZZ5O5OPQ5R'RW^&_G V#2C(CNNMhWZ^[delmpdqq{  }N  |O  O{  | +/*B*BCUW]*^'GZ")500#bcubv  wI  "JJ!%!=!=*&< "> "J KK"9* FG!,,U3#33%43I3I3K &/Z [:>:N:NOa:b 7Z#.+5t+;$; &1B:,Ng/h i"NN+LM262J2JK]_e2f/GZ#*'15'8 8+jk}j~  Q  *RJE!&*&>&>?QSY&Z#%u,,t##E % &"%&H  	m4SVH<LYMaMaMcLdeglll	ms   P9e <He Ae !	e +4e  	e *A+e Bd: %Ae 3Be Ae &e 6e :e	e e		e 	f$e?9f?fr2  c                    t         j                  j                         sd| _        t	        j
                  d       y	 ddlm}m} 	 t         j                  j                         }t         j                  j                  |      }t         j                  j                          | j                          | j                  |dd	d
      \  }}|rd| _        nd| _        t	        j                   d| d| d| j                   d|d       t#        j"                         }	 |||d| j                  || j$                        | _        t	        j                   dt#        j"                         |	z
  dd       d| _        d| _        d| d| d|dd| S # t        $ r d| _        t	        j
                  d       Y yw xY w# t*        $ r4}
d| _        dt-        |
       dt/        j0                          cY d }
~
S d }
~
ww xY w)!zInitialize 5Hz LM model using vllm backend. When enforce_eager is True, CUDA graph
        capture is disabled (required when LoRA training may run in the same process).Fz8CUDA/ROCm is not available. Please check your GPU setup.u<   ❌ CUDA/ROCm is not available. Please check your GPU setup.r   )LLMSamplingParamszfnano-vllm is not installed. Please install it using 'cd acestep/third_parts/nano-vllm && pip install .uj   ❌ nano-vllm is not installed. Please install it using 'cd acestep/third_parts/nano-vllm && pip install .   皙?r   )r   r   r   r      r7   z Initializing 5Hz LM with model: z, enforce_eager: z*, tensor_parallel_size: 1, max_model_len: z, gpu_memory_utilization: r   r   )r  r2  tensor_parallel_sizerB   gpu_memory_utilizationr,  z#5Hz LM initialized successfully in r)  r*  TrX   r   z	
Device: z
GPU Memory Utilization: z
Low GPU Memory Mode: r   r   N)rD   rb   rc   r@   r
   errornanovllmrL  rM  ImportErrorcurrent_deviceget_device_namer\   r_   r   rB   r   r4  r?   r>   rA   r^   r   r   r   )rT   r   r2  rL  rM  rV  device_namerR  r   r<  r   s              r.   r:  z"LLMHandler._initialize_5hz_lm_vllm  sD    zz&&(#(D LLSTQ	@4$	f"ZZ668N**44^DKJJ""$113 ;?:Y:Y%	 ;Z ;7"$7 #%)"%)"KK::,FWXeWf  gQ  RV  Rd  Rd  Qe  e  @V  WZ  [  \  ]J +%&"00'=,,DH KK=diikJ>VWZ=[[cde#'D %DA*ZXcWdd~  @V  WZ  [  [r  sF  rG  H  HM  	@#(D LL  B  C	@N  	f#(D 4SVH<LYMaMaMcLdee	fs0   F EF8 %F54F58	G5)G0*G50G5	cfg_scaleseedsc                    ddl m} | j                  |      \  }}t        |      }| xr
 |
duxs |du}|rdn|}| j	                  |xs ||	|||||||||
|      }| j                  ||| j                  dz
        } |||||||||r|j                  nd      }|dkD  r<| j                  ||||||	      } | g|z  }!| j                  j                  |||!
      }"n| j                  j                  ||      }"g }#|"D ]  }$t        |$d      rAt        |$j                        dkD  r)|#j                  |$j                  d   j                         Pt        |$d      r|#j                  |$j                         xt        |$t               rd|$v r|#j                  |$d          |#j                  t#        |$              |s|#d   S |#S )a  
        Unified vllm generation function supporting both single and batch modes.
        Accepts either a single formatted prompt (str) or a list of formatted prompts (List[str]).
        Returns a single string for single mode, or a list of strings for batch mode.
        r   )rM  Nr   )r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   )
max_tokensr  rY  r   r   r   r   logits_processor_update_stater   r   r   r   r   r   )unconditional_promptsr  text)rT  rM  r#  lenr   r   rB   r  r   r>   generaterL   r  r|   ra  r!  dictr   )%rT   r  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rZ  rM  formatted_prompt_listr   
batch_sizer   effective_sampler_temprP   r]  sampling_paramsformatted_unconditional_promptr`  r  output_textsoutputs%                                        r.   	_run_vllmzLLMHandler._run_vllm  s   < 	, +/*E*EFW*X'x./

 &.!u3Gt3S3tWhptWt(>K !% A A%=%WAW'A+'/#%'-!5/ !B !
  11+-++b0 2 

 )!.12Pe*?*L*Lko	
 s?-1-M-M! /!1! .N .* &D$Dz$Q!hh''%&; ( G hh''(=OG Fvy)c&...AA.E##FNN1$5$:$:;(##FKK0FD)f.>##F6N3##CK0  '/|A@L@r0   formatted_promptc                 j   | j                  |ddd      }| j                  ||	|
||||||d
      }| j                         5  |j                         D ci c]!  \  }}||j	                  | j
                        # }}}| j                  |
|t        | j                  j                  dd            }| j                  |      }|d	kD  r
| j                  |||||d
      }||g}| j                   j                  }d| j                   _        | j                  |ddd      }|| j                   _        |j                         D ci c]!  \  }}||j	                  | j
                        # }}}|d   }|j                  dd      }| j                  ||||||||| j                   j                  xs | j                   j                   d|      } | dd } n|rZ| j#                  |d   |j                  d      |||||| j                   j                  xs | j                   j                   d|
      } nt%        j&                         5   | j                  j(                  di |||dkD  r|nd	|dkD  rdnd||dkD  r|nd|d|cxk  rd	k  rn n|ndt+        |      dkD  r|nd| j                   j                  xs | j                   j                   dd} ddd       ddd       t-         t$        j.                        r| j1                         dk(  r| d   }!n| }!n| d   }!|d	kD  rd   j2                  d   }"n|d   j2                  d   }"|!|"d }!|!j
                  j4                  dk7  r|!j7                         }!| j                   j9                  |!d      }#|#S c c}}w c c}}w # 1 sw Y   xY w# 1 sw Y   xY w)z<Internal helper function for single-item PyTorch generation.r   FTreturn_tensorspadding
truncation
r   r   r   r   r   r   r   r   r   r   r   r7   r\  r   r_  leftr  attention_maskN)batch_input_idsbatch_attention_maskr   r  rY  r   r   r   r  streamerrP   r   r   )
r  ru  r   r  r   r   r   r  rx  rP   r   )r   r  	do_sampler   r   r   r  rx  r!   r8   skip_special_tokensr  )r?   r   _load_model_contextitemsr   rC   r   r%   r>   configr   r   padding_siderJ   _generate_with_cfg_customr  r
  #_generate_with_constrained_decodingrD   inference_moderc  rb  r!  Tensorr   r  typer8   decode)$rT   rm  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   inputsrP   kr-   r   r   ri  batch_textsoriginal_padding_sidebatch_inputs_tokenizedrv  rw  r  r  input_lengthoutput_texts$                                       r.   _run_pt_singlezLLMHandler._run_pt_single:  s   . ##	 $ 
 !% A A%='A+'/#%'- !B !
 %%'7=||~F~tq!adkk**~FF "99 /!1$TXX__6FM : N  $;;<NO3151Q1Q#!%$3%5" 2R 2.  01OP(,(:(:(G(G%28""/)-););#' #	 *< *& 3H""/KaKgKgKi)jKi41a!QTT$++->*>Ki&)j #9"E'='A'ABRTX'Y$ 88$3)=#1 +''9!%!3!3!@!@!cDDVDVDcDc!*? 9  "!A,)BB$[1#)::.>#?#1 +'9!%!3!3!@!@!cDDVDVDcDc!*? C  ))+/dhh// 
 
'53>?K*5/$u','8UQYeD','8S5=N3=NeTX=@AQ=RUV=V)9\`%)%7%7%D%D%gHZHZHgHg!%
G ,W (t gu||,{{}! '
 '#AJM s? 2+>DDQGL!+.44Q7L%lm4 $$-)--/M((//SX/Yc GF *kN ,+W ('sE   N)&N;CN)&N)CN)BNN)N)N&	"N))N2c                    | j                  |      \  }}|r5g }| j                         5  t        |      D ]  \  }}|r|t        |      k  rt	        j
                  ||          t        j                  j                         r#t        j                  j                  ||          ndt        t        j                  d      rJt        j                  j                  j                         r"t        j                  j                  ||          | j                  |||||||||	|
ddddd||||      }|j                  |        	 ddd       |S |d   }| j                  |||||||||	|
|||||||||      S # 1 sw Y   |S xY w)a  
        Unified PyTorch generation function supporting both single and batch modes.
        Accepts either a single formatted prompt (str) or a list of formatted prompts (List[str]).
        Returns a single string for single mode, or a list of strings for batch mode.
        Note: PyTorch backend processes batch items sequentially (doesn't support true batching efficiently).
        rZ   NFTrm  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   )r#  r|  	enumeraterb  rD   manual_seedrb   rc   manual_seed_allrL   rd   rZ   r  r|   )rT   r  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rZ  re  r   rj  irm  r  s                              r.   _run_ptzLLMHandler._run_pt  s   < +/*E*EFW*X'x
 L))++45J+K'A'SZ))%(3 ::224!JJ66uQx@$U^^U;@R@R@_@_@a!II11%(; #'"5"5)9$/"+(7##+=1I3M(7&**/$(%)&*)9 '%!)' #6 #K, !''4A ,L ,F   13""-#+1%='A+'/#%'-' # 
 	
Q ,F  s   DE66F c                 ,    |yd|v rd|v r	d|v rd|v ryy)z+Check if all required metadata are present.FbpmkeyscaletimesignaturedurationTr  )rT   r   s     r.   has_all_metaszLLMHandler.has_all_metas>  s9     M!jM&AoYfFfku  zG  lGr0   metadatac                 b   i }dD ]l  }||v s||   ||   }|dk(  r%|j                  d      r|j                  d      d   }t        |t              r|j	                         rt        |      }|||<   n t        |      dkD  r't        j                  |dd      j                         }nd}d	| d
S )a+  
        Format parsed metadata as CoT text using YAML format (matching training format).

        Args:
            metadata: Dictionary with keys: bpm, caption, duration, keyscale, language, timesignature

        Returns:
            Formatted CoT text: "<think>
{yaml_content}
</think>"
        )r  r   r  r  languager  r  z/4/r   T)allow_unicode	sort_keysr    z<think>
z	
</think>)
endswithr+   r!  r   isdigitr   rb  yamldumpr   )rT   r  	cot_itemskeyvaluecot_yamls         r.   _format_metadata_as_cotz"LLMHandler._format_metadata_as_cotF  s     	ZCh8C=#< /)ennT.B!KK,Q/EeS)emmoJE!&	# [ y>Ayy$$OUUWHH8*J//r0   
infer_typeuse_cot_metasuse_cot_captionuse_cot_languagerf  c                    |d }|xs dj                         j                         }|dvr"d|d}|r|dkD  rg ni |r|dkD  rg ndd|d	i id
S |xr |dkD  }|r|nd}i }d}| j                  |      }d}d}|r|-t        |      D cg c]  }t	        j
                  dd       }}nYt        |      |k  rFt        |      t        |t        |      z
        D cg c]  }t	        j
                  dd       c}z   }n|d| } |dd       |sE|rB|rt        j                  d       nt        j                  d       t        j                         }| j                  ||d      }t        j                  d|        | j                  |||||||	d|| | dd||d|
|d      \  }} t        j                         |z
  }|s|rg ni |rg ndd| d	d|iid
S | j                  |      \  }}|r4t        j                  d|ddt        |j                                       nt        j                  d|ddt        |j                                       nS|rt        j                  d       nt        j                  d       |j                         D !"ci c]  \  }!}"|"	|!|" }}!}"|dk(  rF|r7t        |      D cg c]  }|j!                          }#}|#dg|z  ddd	||d id
S |dddd	||d id
S |rt        j                  d!| d"       nt        j                  d#       t        j                         }$| j#                  |      }%| j%                  |||%      }&t        j                  d$|&         |d%d&| d"       |rU|&g|z  }'	 | j&                  d'k(  r!| j)                  |'||||||	|
||d(|||%|)      }(nP| j&                  d*k(  r!| j+                  |'||||||	|
||d(|||%|)      }(n | j-                  |'||||||	|
||d(|||%|)      }(g }*g }#|(D ]F  }+| j                  |+      \  }},|*j5                  |,       |#j5                  |j!                                H t        j                         |$z
  }|*D -cg c]#  }-|-rt        |-j7                  d-            dz
  nd% }.}-t        j                  d.|dd/|.        ||z   }/|#|*dd|||/d,|.t9        |.      d0d
S | j                  |&||||||	|dddd(|||%d1|
|d      \  }0} |0s||z   }/|dd| d	|||/d,id
S t        j                         |$z
  }| j                  |0      \  }}|rt        |j7                  d-            dz
  nd}1t        j                  d2|dd3|1 d4       ||z   }/||dd|||/d,|1d5d
S c c}w c c}w c c}"}!w c c}w # t.        $ r;})d+t1        |)       }t        j2                  |       g g d|d	|d|d,id
cY d})~)S d})~)ww xY wc c}-w )6a  Two-phase LM generation: CoT generation followed by audio codes generation.

        - infer_type='dit': Phase 1 only - generate CoT and return metas (no audio codes)
        - infer_type='llm_dit': Phase 1 + Phase 2 - generate CoT then audio codes

        Args:
            target_duration: Target duration in seconds for codes generation constraint.
                            5 codes = 1 second. If specified, blocks EOS until target reached.
            user_metadata: User-provided metadata fields (e.g. bpm/duration/keyscale/timesignature).
                           If specified, constrained decoding will inject these values directly.
            use_cot_caption: Whether to generate caption in CoT (default True).
            use_cot_language: Whether to generate language in CoT (default True).
            batch_size: Optional batch size for batch generation. If None or 1, returns single result.
                       If > 1, returns batch results (lists).
            seeds: Optional list of seeds for batch generation (for reproducibility).
                  Only used when batch_size > 1. TODO: not used yet

        Returns:
            Dictionary containing:
                - metadata: Dict or List[Dict] - Generated metadata
                - audio_codes: str or List[str] - Generated audio codes
                - success: bool - Whether generation succeeded
                - error: Optional[str] - Error message if failed
                - extra_outputs: Dict with time_costs and other info
        Nc                       y Nr  )argskwargss     r.   progressz9LLMHandler.generate_with_stop_condition.<locals>.progress  s    r0   r    >   ditllm_ditzinvalid infer_type: z (expected 'dit' or 'llm_dit')r   F
time_costs)r  audio_codesrD  rS  extra_outputsr   r   l    rO  z8Phase 1: Generating CoT metadata (once for all items)...z>Batch Phase 1: Generating CoT metadata (once for all items)...z#Phase 1: Generating CoT metadata...r   )r   z/generate_with_stop_condition: formatted_prompt=T)r  rY  r   r   r   r   r   r   r   r   r   r   r   r   rm  cfgr   r   r   phase1_timezBatch Phase 1 completed in r)  zs. Generated metadata: zPhase 1 completed in zABatch Phase 1: Using user-provided metadata (skipping generation)z;Phase 1: Using user-provided metadata (skipping generation)r  )r  
total_timez*Batch Phase 2: Generating audio codes for z	 items...z"Phase 2: Generating audio codes...z8generate_with_stop_condition: formatted_prompt_with_cot=r   z$Phase 2: Generating audio codes for rX   r   )r  r  rY  r   r   r   r   r   r   r   r   r   r   r   rZ  r/  z!Error in batch codes generation: )r  phase2_timer  <|audio_code_zBatch Phase 2 completed in zs. Generated codes: )r  codes_countstotal_codes)r  rY  r   r   r   r   r   r   r   r   r   r   r   r   zPhase 2 completed in zs. Generated z audio codes)r  codes_count)r   rK   r  r  randomrandintrb  r"  r
   r   r4  r   generate_from_formatted_promptparse_lm_outputkeysr}  copyr  r   rA   rl  _run_mlxr  r^   r   rS  r|   r+   sum)2rT   r   r   r  r  rY  r   r   r   r   r   r   r   r   r  r  r  rf  rZ  r  	error_msgr   actual_batch_sizer  r  r  r  r  rH  phase1_startrm  cot_output_textstatusr  r-   metadata_listphase2_startr   formatted_prompt_with_cotr  codes_outputsr   audio_codes_listr  audio_codes_itemr   r  r  codes_output_textr  s2                                                     r.   generate_with_stop_conditionz'LLMHandler.generate_with_stop_conditionc  s   ^  !&B--/557
//.zn<Z[I#-*q.Br&0Z!^r" "".!3  0*q.*2J **=9 }?DEV?WX?W!95?WXU//UUSdgjkpgqSqMr&sMrv~~a'CMr&ss001 	PR\]AB99;L  $::7F]b:cKKIJZI[\]&*&I&I!1#.!*'6""*<'+%2(7$7)9%9#'(-&$" *B+E"&+ 'J '#OV0 ))+4K"&.B)12r$#&2]K4P%Q  ..?KHa9+c9JJabfgogtgtgvbwaxyz3K3DD[\`aiananap\q[rst _`YZ)6)<)<)>P)>A!-1)>HP :?@Q:R S:RQ:R S -$&4*;#;#!$+6*5'&  !)#%#!$+6*5'&  KKDEVDWW`abKK<=yy{ //9 %)$H$HRXZb$c!NOhNijk<=N<OyYZ!: ;>O OF##v-$(NN*;$/"+(7##+=1I3M(7)0 '%!)# %3 %M" %%.$(MM*;$/"+(7##+=1I3M(7)0 '%!)# %2 %M$ %)LL*;$/"+(7##+=1I3M(7)0 '%!)# %1 %MF  "M,&*&:&:;&G## ''(89$$X]]_5  -
 ))+4K `pp_oV[UCO <=APQQ_oLpKK5k#5FFZ[gZhij${2J)/ (3'2&0#
 %1#&|#4" " )-(K(K!:#.!*'6""*<'6%)$(%)(/&$ (" *B+E"'+ )L )%v0 %(;6
 (#%$#$+6+6*4'&  ))+4K "112CDNA{IT#k//@AAEZ[KKK/C/@k]Zfgh${2J$* (3'2&0#
 $/" c Y&sD Q
 !TL  ?AxH	Y' "#%$&$+6+.*5'& 8 qsC   VV
VV8VB V$ 1(W+$	W(-0W#W(#W(r   c                     | j                   t        d      |r,| j                  |      }|dk(  r|r
d| d| d}nd| d}n|}n	d| d| d}| j                   j                  ddt         d	d
d|d
gdd      S )aA  
        Build the chat-formatted prompt for 5Hz LM from caption/lyrics.
        Raises a ValueError if the tokenizer is not initialized.

        Args:
            caption: Caption text
            lyrics: Lyrics text
            is_negative_prompt: If True, builds unconditional prompt for CFG
            generation_phase: "cot" or "codes" - affects unconditional prompt format
            negative_prompt: Negative prompt for CFG (used when is_negative_prompt=True)

        Example:
            prompt = handler.build_formatted_prompt("calm piano", "hello world")
        :LLM tokenizer is not initialized. Call initialize() first.r   
# Caption



# Lyric

z# Lyric
system# Instruction


rolecontentuserFTtokenizeadd_generation_prompt)r?   
ValueErrorr   apply_chat_templater   )rT   r   r   r   r   r   has_negative_promptprompts           r.   r   z!LLMHandler.build_formatted_prompt  s     %YZZ #'"F"F"W5(&*?*;=PRSF  )3F ! #7)=CF!!55!@V?WW[.\]F3 "& 6 
 	
r0   c                    | j                   t        d      |r| j                  |      }d}|r|}n|}n|}|}d| d| d}	| j                   j                  ddt         dd	d
|	d	d|d	gdd      }
|
j                  d      s|
dz  }
|
S )a  
        Build the chat-formatted prompt for codes generation phase with pre-generated CoT.

        Args:
            caption: Caption text
            lyrics: Lyrics text
            cot_text: Pre-generated CoT text (e.g., "<think>\nbpm: 120\n...\n</think>")
            is_negative_prompt: If True, uses empty CoT for CFG unconditional prompt
            negative_prompt: Negative prompt for CFG (used when is_negative_prompt=True)

        Returns:
            Formatted prompt string

        Example:
            cot = "<think>\nbpm: 120\ncaption: calm piano\n...\n</think>"
            prompt = handler.build_formatted_prompt_with_cot("calm piano", "hello", cot)
        r  z<think>
</think>r  r  r  r  r  r  r  r  	assistantFr  )r?   r  r   r  r   r  )rT   r   r   r   r   r   r  cot_for_promptcaption_for_promptuser_prompt	formatteds              r.   r   z*LLMHandler.build_formatted_prompt_with_cot  s    $ %YZZ #'"F"F"W 1N"%4" &-" &N!( $$6#7}VHBO &&::!@V?WW[.\]K8$@
 "' ; 
	 !!$'Ir0   r  c                     | j                   t        d      |r|r|j                         r|nd}n|}| j                   j                  ddt         ddd|dgdd	
      S )a  
        Build the chat-formatted prompt for audio understanding from codes.

        This is the reverse of generation: given audio codes, generate metadata and lyrics.

        Args:
            audio_codes: Audio code string (e.g., "<|audio_code_123|><|audio_code_456|>...")
            is_negative_prompt: If True, builds unconditional prompt for CFG
            negative_prompt: Negative prompt for CFG (used when is_negative_prompt=True)

        Returns:
            Formatted prompt string

        Example:
            codes = "<|audio_code_18953|><|audio_code_13833|>..."
            prompt = handler.build_formatted_prompt_for_understanding(codes)
        r  r    r  r  r  r  r  FTr  )r?   r  r   r  r   )rT   r  r   r   user_contents        r.   (build_formatted_prompt_for_understandingz3LLMHandler.build_formatted_prompt_for_understanding=  s    . %YZZ .=/BWBWBY?_aL&L!!55 %!01R0SSWX
 #+	 "& 6 
 	
r0   c                    t        | dd      si dfS |r|j                         si dfS t        j                  dt	        |       d       | j                  |      }t        d|        | j                  |||||dddddd	d
d
d||d      \  }	}
|	si |
fS | j                  |	      \  }}| j                  |	      }|r||d<   t        j                  dt	        |       d       |rKt        j                  dt        |j                                       t        j                  d|	dd  d       ddj                  |j                                }||fS )a	  
        Understand audio codes and generate metadata + lyrics.

        This is the reverse of the normal generation flow:
        - Input: Audio codes
        - Output: Metadata (bpm, caption, duration, etc.) + Lyrics

        Note: cfg_scale and negative_prompt are not supported in understand mode.

        Args:
            audio_codes: String of audio code tokens (e.g., "<|audio_code_123|><|audio_code_456|>...")
            temperature: Sampling temperature for generation
            top_k: Top-K sampling (None = disabled)
            top_p: Top-P (nucleus) sampling (None = disabled)
            repetition_penalty: Repetition penalty (1.0 = no penalty)
            use_constrained_decoding: Whether to use FSM-based constrained decoding for metadata
            constrained_decoding_debug: Whether to enable debug logging for constrained decoding

        Returns:
            Tuple of (metadata_dict, status_message)
            metadata_dict contains:
                - bpm: int or str
                - caption: str
                - duration: int or str
                - keyscale: str
                - language: str
                - timesignature: str
                - lyrics: str (extracted from output after </think>)

        Example:
            codes = "<|audio_code_18953|><|audio_code_13833|>..."
            metadata, status = handler.understand_audio_from_codes(codes)
            print(metadata['caption'])  # "A cinematic orchestral piece..."
            print(metadata['lyrics'])   # "[Intro: ...]\n..."
        r@   F7   ❌ 5Hz LM not initialized. Please initialize it first.u<   ❌ No audio codes provided. Please paste audio codes first.z#Understanding audio codes (length: z chars)zformatted_prompt: N
understandr    r  r   r   r   r   r   r   r   r   r   r   r   r  r   z#Understanding completed. Generated z metadata fieldsGenerated metadata: Output text preview:    ...u;   ✅ Understanding completed successfully
Generated fields: , )r%   r   r
   r   rb  r  printr  r  _extract_lyrics_from_outputr   r"  r  rq   )rT   r  r  r   r   r   r   r   rm  r  r  r  rH  r   r   s                  r.   understand_audio_from_codesz&LLMHandler.understand_audio_from_codesm  s   Z t.6PPP+"3"3"5UUU9#k:J9K7ST  HHU"#3"456 #AA-*&8#'!% %!&$$0 &>'A#' B 
V, v: **;7! 11+>!'HX9#h-HXYZ%LL/X]]_0E/FGHLL0Tc1B0C3GHSTXT]T]^f^k^k^mTnSop
##r0   r  c                    ddl }d}|j                  ||      }|sy||j                         d j                         }|syd}|j	                  |d||j
                        }|j	                  dd|      }|j                         S )aM  
        Extract lyrics section from LLM output.

        The lyrics appear after the </think> tag and typically start with "# Lyric"
        or directly with lyric content.

        Args:
            output_text: Full LLM output text

        Returns:
            Extracted lyrics string, or empty string if no lyrics found
        r   Nr3   r    z^#\s*Lyri[c|cs]?\s*\n)flagsz<\|im_end\|>\s*$)researchendr   sub
IGNORECASE)rT   r  r  think_end_patternmatchafter_thinklyric_header_patterns          r.   r  z&LLMHandler._extract_lyrics_from_output  s     	 (		+[9 "%))+,/557  8ff12{"--fX ff0"kB  ""r0   queryinstrumentalc                     | j                   t        d      |rdnd}|r|r|j                         r|nd}n| d| }| j                   j                  ddt         dd	d
|d	gdd      S )a  
        Build the chat-formatted prompt for inspiration/simple mode.

        This generates a complete sample (caption, lyrics, metadata) from a user's
        natural language music description query.

        Args:
            query: User's natural language music description
            instrumental: Whether to generate instrumental music (no vocals)
            is_negative_prompt: If True, builds unconditional prompt for CFG
            negative_prompt: Negative prompt for CFG (used when is_negative_prompt=True)

        Returns:
            Formatted prompt string

        Example:
            query = "a soft Bengali love song for a quiet evening"
            prompt = handler.build_formatted_prompt_for_inspiration(query, instrumental=False)
        r  r;   falser    z

instrumental: r  r  r  r  r  FTr  )r?   r  r   r  r   )rT   r
  r  r   r   instrumental_strr  s          r.   &build_formatted_prompt_for_inspirationz1LLMHandler.build_formatted_prompt_for_inspiration  s    4 %YZZ &26w.=/BWBWBY?_aL $W$67G6HIL!!55 %!01P0QQUV
 #+	 "& 6 
 	
r0   vocal_languagec
                 l   t        | dd      si dfS |r|j                         sd}t        j                  d|dd  d| d	| d
       | j	                  ||      }
t        j
                  d|
        d}d}|ri|j                         rY|j                         j                         dk7  r8d|j                         i}t        j                  d|j                                 | j                  |
||||d|ddddddd||	d      \  }}|si |fS | j                  |      \  }}| j                  |      }|r||d<   n|rd|d<   ||d<   t        j                  d| d       |	rKt        j
                  dt        |j                                       t        j
                  d|dd  d       d| }||fS )a  
        Create a complete music sample from a user's natural language query.

        This is the "Simple Mode" / "Inspiration Mode" feature that generates:
        - Metadata (bpm, caption, duration, keyscale, language, timesignature)
        - Lyrics (unless instrumental=True)

        Args:
            query: User's natural language music description
            instrumental: Whether to generate instrumental music (no vocals)
            vocal_language: Allowed vocal language for constrained decoding (e.g., "en", "zh").
                           If provided and not "unknown", it will be used.
            temperature: Sampling temperature for generation (0.0-2.0)
            top_k: Top-K sampling (None = disabled)
            top_p: Top-P (nucleus) sampling (None = disabled)
            repetition_penalty: Repetition penalty (1.0 = no penalty)
            use_constrained_decoding: Whether to use FSM-based constrained decoding
            constrained_decoding_debug: Whether to enable debug logging

        Returns:
            Tuple of (metadata_dict, status_message)
            metadata_dict contains:
                - bpm: int or str
                - caption: str
                - duration: int or str
                - keyscale: str
                - language: str
                - timesignature: str
                - lyrics: str (extracted from output after </think>)
                - instrumental: bool (echoed back)

        Example:
            query = "a soft Bengali love song for a quiet evening"
            metadata, status = handler.create_sample_from_query(query, instrumental=False, vocal_language="bn")
            print(metadata['caption'])  # "A gentle romantic acoustic pop ballad..."
            print(metadata['lyrics'])   # "[Intro: ...]\n..."
        r@   Fr  r   zCreating sample from query: Nd   z... (instrumental=z, vocal_language=r+  )r
  r  z"Formatted prompt for inspiration: unknownr  zUsing user-specified language: r  r    r  r  r   [Instrumental]r  z'Sample created successfully. Generated  fieldsr  r  ,  r  u2   ✅ Sample created successfully
Generated fields: )r%   r   r
   r   r  r   rK   r  r  r  r"  r  )rT   r
  r  r  r  r   r   r   r   r   rm  r   r   r  r  r  rH  r   r   s                      r.   create_sample_from_queryz#LLMHandler.create_sample_from_query-  s&   b t.6PPPEKKM#E25#;-?QR^Q__pq  qA  AB  C  	D  FF% G 
 	9:J9KLM n2249M9M9O9U9U9W[d9d')=)=)?@MKK9.:N:N:P9QRS
 #AA-*&8#'!. %!&$$0 &>'A#% B 
V* v: **;7! 11+>!'HX!1HX $0 =hZwOP%LL/X]]_0E/FGHLL0Tc1B0C3GHJ8*U
##r0   c                     | j                   t        d      |r|r|j                         r|nd}nd| d| }| j                   j                  ddt         ddd	|dgd
d      S )a  
        Build the chat-formatted prompt for format/rewrite mode.

        This formats user-provided caption and lyrics into a more detailed and specific
        musical description with metadata.

        Args:
            caption: User's caption/description of the music
            lyrics: User's lyrics
            is_negative_prompt: If True, builds unconditional prompt for CFG
            negative_prompt: Negative prompt for CFG (used when is_negative_prompt=True)

        Returns:
            Formatted prompt string

        Example:
            caption = "Latin pop, reggaeton, flamenco-pop"
            lyrics = "[Verse 1]\nTengo un nudo..."
            prompt = handler.build_formatted_prompt_for_format(caption, lyrics)
        r  r    r  r  r  r  r  r  r  FTr  )r?   r  r   r  r   )rT   r   r   r   r   r  s         r.   !build_formatted_prompt_for_formatz,LLMHandler.build_formatted_prompt_for_format  s    6 %YZZ.=/BWBWBY?_aL )	vhGL!!55 %!01O0PPTU
 #+	 "& 6 
 	
r0   c
                    t        | dd      si dfS |r|j                         sd}|r|j                         sd}t        j                  d|dd  d	t	        |              | j                  ||
      }
t        j                  d|
        d}|ri }|j                  d      	 t        |d         }|dkD  r||d<   |j                  d      	 t        |d         }|dkD  r||d<   |j                  d      r|d   |d<   |j                  d      r|d   |d<   |j                  d      r|d   |d<   |sd}nt        j                  d|        | j                  |
||||d|d|r|j                  d      dundddddd||	d      \  }}|si |fS | j                  |      \  }}| j                  |      }|r||d<   n||d<   t        j                  d| d       |	rKt        j                  dt        |j                                       t        j                  d|dd  d       ddj!                  |j                                }||fS # t        t        f$ r Y w xY w# t        t        f$ r Y w xY w) as  
        Format user-provided caption and lyrics into structured music metadata.

        This is the "Format" feature that takes user input and generates:
        - Enhanced caption with detailed music description
        - Metadata (bpm, duration, keyscale, language, timesignature)
        - Formatted lyrics (preserved from input)

        Note: cfg_scale and negative_prompt are not supported in format mode.

        Args:
            caption: User's caption/description (e.g., "Latin pop, reggaeton")
            lyrics: User's lyrics with structure tags
            user_metadata: Optional dict with user-provided metadata to constrain decoding.
                          Supported keys: bpm, duration, keyscale, timesignature, language
            temperature: Sampling temperature for generation (0.0-2.0)
            top_k: Top-K sampling (None = disabled)
            top_p: Top-P (nucleus) sampling (None = disabled)
            repetition_penalty: Repetition penalty (1.0 = no penalty)
            use_constrained_decoding: Whether to use FSM-based constrained decoding
            constrained_decoding_debug: Whether to enable debug logging

        Returns:
            Tuple of (metadata_dict, status_message)
            metadata_dict contains:
                - bpm: int or str
                - caption: str (enhanced)
                - duration: int or str
                - keyscale: str
                - language: str
                - timesignature: str
                - lyrics: str (from input, possibly formatted)

        Example:
            caption = "Latin pop, reggaeton, flamenco-pop"
            lyrics = "[Verse 1]\nTengo un nudo en la garganta..."
            metadata, status = handler.format_sample_from_input(caption, lyrics)
            print(metadata['caption'])  # "A dramatic and powerful Latin pop track..."
            print(metadata['bpm'])      # 100
        r@   Fr  r   r  z&Formatting sample from input: caption=N2   z..., lyrics length=)r   r   zFormatted prompt for format: r  r   r  r  r  r  z*Using user-provided metadata constraints: r  r    r  r  r   z)Format completed successfully. Generated r  r  r  r  r  u4   ✅ Format completed successfully
Generated fields: r  )r%   r   r
   r   rb  r  r   rJ   r   r  	TypeErrorr  r  r  r"  r  rq   )rT   r   r   r   r  r   r   r   r   r   rm  constrained_metadatabpm_valdur_valr  r  r  rH  formatted_lyricsr   s                       r.   format_sample_from_inputz#LLMHandler.format_sample_from_input  s   h t.6PPPgmmo%GV\\^%F<WSb\NJ]^abh^i]jkl  AA B 
 	45E4FGH  $#%   '3!-"67G{6=,U3   ,8!-
";<G{;B,Z8   ,3@3L$Z0  18Eo8V$_5  ,3@3L$Z0 ('+$HI]H^_`
 #AA-*&8#'!5 %Ui!5!9!9*!ET!Qot$$0 &>'A#% B 
V* v: **;7!  ;;KH!1HX "(HX?zQR%LL/X]]_0E/FGHLL0Tc1B0C3GHLTYYW_WdWdWfMgLhi
##E #I.  #I. s$   $I I# I I #I65I6r  c                    t        | dd      sy| j                  dk(  r| j                  | j                  y| j                  | j                  y|xs i }|j                  dd	      }|j                  d
d      }|j                  dd      }|j                  d      }	|j                  d      }
|j                  dd      }|j                  d      }|j                  d      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }	 | j                  dk(  r4| j                  |||||	|
|||||||||||||      }|dt        |       fS | j                  dk(  r4| j                  |||||	|
|||||||||||||      }|dt        |       fS | j                  |||||	|
|||||||||||||      }|d t        |       fS # t        $ r4}d!dl} |j                         }t        j                  d"t        |      j                    d#| d$|        | j                  dk(  r_	 d!d%lm}  |        n# t&        $ r Y nw xY w	 t)        | j                  d&      r| j                  j+                          n# t        $ r Y nw xY wt,        j.                  j1                         r=t,        j.                  j3                          t,        j.                  j5                          nt)        t,        d'      r[t,        j6                  j1                         r=t,        j6                  j3                          t,        j6                  j5                          n~t)        t,        j8                  d(      rdt,        j8                  j:                  j1                         r<t,        j:                  j3                          t,        j:                  j5                          dd)t        |      j                    d#|xs |j=                         d*    fcY d}~S d}~ww xY w)+aw  
        Generate raw LM text output from a pre-built formatted prompt.

        Args:
            formatted_prompt: Prompt that is already formatted by `build_formatted_prompt`.
            cfg: Optional dict supporting keys:
                - temperature (float)
                - cfg_scale (float)
                - negative_prompt (str) used when cfg_scale > 1
                - top_k (int), top_p (float), repetition_penalty (float)
                - target_duration (float): Target duration in seconds for codes generation
                - generation_phase (str): "cot" or "codes" for phase-aware CFG
            use_constrained_decoding: Whether to use FSM-based constrained decoding
            constrained_decoding_debug: Whether to enable debug logging for constrained decoding
            stop_at_reasoning: If True, stop generation immediately after </think> tag (no audio codes)

        Returns:
            (output_text, status_message)

        Example:
            prompt = handler.build_formatted_prompt(caption, lyric)
            text, status = handler.generate_from_formatted_prompt(prompt, {"temperature": 0.7})
        r@   F)r    r  r/  N)r    u-   ❌ 5Hz LM is missing MLX model or tokenizer.)r    u)   ❌ 5Hz LM is missing model or tokenizer.r  g333333?rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r    r   r   rX   )r  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   u+   ✅ Generated successfully (vllm) | length=u*   ✅ Generated successfully (mlx) | length=u)   ✅ Generated successfully (pt) | length=r   z)Error in generate_from_formatted_prompt: : r  )reset_contextrY   r]   rZ   u,   ❌ Error generating from formatted prompt: r   )r%   rA   rR   r?   r>   rJ   rl  rb  r  r  r^   r   r   r
   rS  r  __name__nanovllm.utils.contextr$  rU  rL   rY   rD   rb   rc   r\   r[   r]   rd   rZ   
splitlines)rT   rm  r  r   r   r   r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r  r   r   error_detailr$  s                             r.   r  z)LLMHandler.generate_from_formatted_promptm  sJ   > t.6Pu$&$*<*<*DJXX!3!3!;BiRggmS1GGK-	''"3_E   WW%93?''"340ww~u57ggmU377#5u='')R(2&77:r*h	6)"nn&6 +'$3'9-E/I$3"/&7 +!-"/%5#!%' - * #&QRUVaRbQc$ddd!!U*"mm&6 +'$3'9-E/I$3"/&7 +!-"/%5#!%' , * #&PQTU`QaPb$ccc ,,"2'# /#5)A+E /+"3')+!1!' ' K* "KCP[L\K] ^^^ 	/9//1LLLDT!WEUEUDVVXYZX[[]^j]klm 6)D!O" txx1(   zz&&(

&&(

&&(&599+A+A+C		%%'		%%'/ENN4F4F4S4S4U		%%'		%%'Ed1gFVFVEWWYZ[Z|_k_v_v_xy{_|Y}~~~=	sv   !AH $AH '3H 
Q%AQ;J	Q		JQJQ0K
	Q
	KQKE8QQQr  ru  r   rx  c           
         | j                   }| j                  }|j                         }||j                         }nt        j                  |      }d|i}d}t        |d      xr t        |j                  dd      }| j                  j                  }||}| j                  |      }t        j                         5  t        t        |      dd| j                        D ]L  }| j                  |||||      }|j                   ddd	ddf   }|
	 |
||      }|D ]  } |||      } | j#                  ||      }| j%                  ||      }| j'                  ||      }| j)                  |
|       | j+                  |||      }|j-                  d
      }t        j.                  ||gd
      }t        j.                  |t        j0                  |j2                  d   d
f||j4                        gd
      }||d<   |rt        |d      r|j6                  }|	|	j9                  |       |sM n ddd       |	|	j;                          |S # 1 sw Y   xY w)z
        Custom generation loop with constrained decoding support (non-CFG).
        This allows us to call update_state() after each token generation.
        Nru  generation_configr  TzLLM Constrained Decodingtokendescunitdisabler   r   r   r   rC   rF   r  )r>   rC   r   rD   	ones_likerL   r%   r*  r?   r
  r   r  r   r  rN   r  r   r   r   r  r  r  	unsqueezecatonesr  rF   r  putr  )rT   r  ru  r   r  r   r   r   r  rx  rP   r  rC   r  	attn_maskr  r  r  r
  r   stepr  next_token_logits	processornext_tokensshould_stopnext_tokens_unsqueezeds                              r.   r  z.LLMHandler._generate_with_constrained_decoding	  s   "  ")%&,,.I	2I ))4 E#67oGED[D[]hjn<o	 ))66'L  778JK!!#U>29SZakok|k|}},,UM<Q`bkl %,NN1b!8$<! )4(=mM^(_% "2I(1-AR(S% "2 %)$<$<=NPU$V!$($<$<=NPU$V! #112C[Q 889NP[\ #33K|\ *5)>)>q)A& %		=:P*QWX Y!IIy%**iooa>PRS=T]cktkzkz2{&|  CD  E	1:-. 2C!D&-&=&=O 'LL!78W ~ $\ LLNc $#s   ;E.I
+I

Irv  rw  c           
         | j                   }| j                  }|j                  d   dz  }d}|}|j                         }||j                         }nt	        j
                  |      }i }|||d<   d}t        |d      xr t        |j                  dd      }| j                  j                  }||	}| j                  |      }t	        j                         5  t        t        |      dd	| j                  
      D ]  }| j!                  |||||      }|j"                  dddddf   }||||z    }||||z    }|j%                         ||j%                         |j%                         z
  z  z   }|||||z    } |||      }||||z    }|D ]  } |||      } | j'                  ||      }| j)                  ||      }| j+                  ||      } | j-                  ||        | j/                  | ||	      }!| j1                  d      }"t	        j2                  ||"j5                  dd      gd      }t	        j2                  |t	        j6                  |dz  df||j8                        gd      }||d<   |rt        |d      r|j:                  }|
|
j=                  |"       |!s n ddd       |
|
j?                          |S # 1 sw Y   xY w)a  
        Custom CFG generation loop that:
        1. Processes both conditional and unconditional sequences in parallel
        2. Applies CFG formula to logits
        3. Samples tokens only for conditional sequences
        4. Applies the same sampled tokens to both conditional and unconditional sequences
        5. Optionally applies constrained decoding via FSM-based logits processor

        Batch format: [cond_input, uncond_input]
        r   r!   Nru  r*  r  TzLLM CFG Generationr+  r,  r   r   r   r0  r  ) r>   rC   r  r   rD   r1  rL   r%   r*  r?   r
  r   r  r   r  rN   r  r   r   r   r   r  r  r  r2  r3  repeatr4  rF   r  r5  r  )#rT   rv  rw  r   r  rY  r   r   r   r  rx  rP   r  rC   rf  cond_start_idxuncond_start_idxr  ru  r  r  r  r
  r   r7  r  r8  cond_logitsuncond_logits
cfg_logitscurrent_input_idsr9  r:  r;  r<  s#                                      r.   r  z$LLMHandler._generate_with_cfg_customn	  sA   0 $**1-2
% (--/+1779N"___=N +-;L)* E#67oGED[D[]hjn<o	 ))66'L  778JK!!#U>29MT[eievevww,,UM<Q`bkl %,NN1b!8$<! 0~j?XY 12BCST^C^ _ +002Y+BSBSBUXeXkXkXmBm5nn
 )4(5n^T^E^(_%!67H*!UJ %2.PZAZ$[!!1I!*+<j!IJ "2 "55j%H
!55j%H
 #11*kJ 889NP[\
 #33K|\ *5)>)>q)A& %		=:P:W:WXY[\:]*^de f!&NEJJ
STVWGXago}  pD  pD  =E  ,F  LM  "N1?-. 2C!D&-&=&=O 'LL!78 s x $x LLN C $#s   G	K !K  K	c                    |j                  d      d   }t        j                  d|        i d}ddl}d}|j	                  ||      }|rdj                  |      }g d}d}|D ]B  }	|j                  |	||j                        }
|
s#|
j                  d      j                         } n |s*d	|v r|j                  d	      d   n|}|j                         }|r|j                  d
      }dg fd}|D ]  }|j                         j                  d      r#|r|d   j                         sud|v rq |        |j                  dd      }t        |      dk(  sd|d   j                         j                         |d   }|j                         sj                  |       |j                  d      s|j                  d      sЉsӉj                  |         |        |fS )a  
        Parse LM output to extract metadata and audio codes.

        Expected format:
        <think>
        bpm: 73
        caption: A calm piano melody
        duration: 273
        genres: Chinese folk
        keyscale: G major
        language: en
        timesignature: 4
        </think>

        <|audio_code_56535|><|audio_code_62918|>...

        Returns:
            Tuple of (metadata_dict, audio_codes_string)
        r3   r   zDebug output text: r    Nz<\|audio_code_\d+\|>)<think>(.*?)</think>rF  z<reasoning>(.*?)</reasoning>r   r  r  c                  .   rډrdj                        } dk(  r	 t        | j                               d<   ndk(  rt        j                  |       d<   ndk(  r	 t        | j                               d<   ncdk(  r| j                         d<   nJdk(  r| j                         d<   n1dk(  r| j                         d<   ndk(  r| j                         d<   d	g y	#  | j                         d<   Y xY w#  | j                         d<   Y 6xY w)
z Save the accumulated field valuer  r  r   r  genresr  r  r  N)rq   r   r   r   postprocess_caption)r  current_keycurrent_value_linesr  s    r.   save_current_fieldz6LLMHandler.parse_lm_output.<locals>.save_current_field#
  s    #6 II&9:E"e+<.1%++-.@HUO %	1.P.d.dej.k+$
2A36u{{}3EHZ0 %0-2[[]*$
2/4{{},$
2/4{{},$749KKM1"&(#)<.3kkmHUOA38;;=HZ0s   C# C= #C:=D<:r!    	)r+   r
   r   r  findallrq   r  DOTALLgroupr   r'   isspacerb  rK   r|   )rT   r  debug_output_textr  r  code_patterncode_matchesreasoning_patternsreasoning_textpatternr  lines_before_codeslinesrL  linepartsfirst_valuerJ  rK  r  s                    @@@r.   r  zLLMHandler.parse_lm_output	  s   ( (--j9!<*+<*=>? /zz,<'',/K
 )GIIg{BII>E!&Q!5!5!7	 * JY]hJh!2!2?!CA!Fny/557N "((.EK"$)@ ::<**3/ Q 1cTk&( !JJsA.E5zQ&+Ahnn&6&<&<&>&+Ah&,,./66{C__S)T__T-B"+2248+ 0  $$r0   c                  8    	 ddl m}  ddl}y# t        $ r Y yw xY w)z4Check if MLX framework is available (Apple Silicon).r   NTF)mlx.corecoremlx_lmrU  )mxrc  s     r.   r7  zLLMHandler._is_mlx_availablec
  s"    	! 		s   
 	c                    	 ddl m} ddlm} t	        j
                  d|        t        j                         }	  ||      \  | _        }|j9                  | j                  j;                                || _        t        j                         |z
  }t	        j
                  d|dd       d| _        d| _         d| d}d|fS # t        $ r}t	        j
                  d| d       ddl	}ddl
m} ddlm}	m}
m}m}  ||      } |
|      }|j                  t!        |d	z              }|st#        d
|       |i }|D ]"  }|j%                  |j                  |             $ t'        t)        |            }|j+                  d      sBt	        j
                  d       |j-                         D ci c]  \  }}d| | nc c}}w }}} ||      \  }}|j/                  |      } ||      }t1        |d      r|j3                  |      }|j5                  t7        |j-                               d       |j9                  |j;                                |j9                          || _        Y d}~(d}~ww xY w# t        $ rJ}ddl!} |jD                         }t	        jF                  d| d|        ddt!        |       fcY d}~S d}~ww xY w)z
        Load the 5Hz LM model using mlx-lm for native Apple Silicon acceleration.

        Args:
            model_path: Path to the HuggingFace model directory

        Returns:
            Tuple of (success, status_message)
        r   N)loadzLoading MLX model from zStandard MLX load failed (z-), retrying with 'model.' prefix remapping...)Path)
load_modelload_configload_tokenizer_get_classeszmodel*.safetensorszNo safetensors found in zmodel.z;Adding 'model.' prefix to weight keys for MLX compatibility)r~  sanitizeT)strictz!MLX model loaded successfully in r)  sr/  r   z>
Backend: MLX (Apple Silicon native)
Device: Apple Silicon GPUzFailed to load MLX model: r  Fu   ❌ MLX load failed: )$ra  rb  mlx_lm.utilsrf  r
   r   r4  rR   r^   globpathlibrg  rh  ri  rj  rk  r   FileNotFoundErrorupdatenextiterr'   r}  	from_dictrL   rl  load_weightsr"  r   
parametersrS   rA   r@   r   r   rj   )rT   r   rd  mlx_loadr<  rH  	first_err_globrg  rh  ri  rj  rk  _model_pathr~  weight_filesweightswf
sample_keyr  r-   model_classmodel_args_class
model_argsr  	load_timer   r   r   r(  s                                 r.   r8  zLLMHandler._load_mlx_modelm
  s   L	;!5KK1*>?J-(%-j%9"\ GGDOO..01#-D 		j0IKK;Ic?!LM$D#'D $ &,-  ##y  +( 0 <A A %(^^":.$[1  %zz#k<P.P*QR#+.Fzl,STZcc&BNN2772;/ ' "$w-0
!,,X6KK ]^;B==?K?41as|Q?KGK 1=F0K---77?
#J/5*-#nnW5G""4#8"F((*+

"'W+(|  	;/9//1LNN7s"\NKL1#a&:::		;sU   8I; C A8I; 
I8C#I31GB,I3-I; 3I88I; ;	K?K	K	Kc                     ddl m} 	 ddlm}  || j                        S # t
        t        f$ r6 	 | j                  j                         cY S # t        $ r t        d      w xY ww xY w)z$Create a KV cache for the MLX model.r   Nmake_prompt_cachez;Cannot create MLX KV cache. Ensure mlx-lm version >= 0.20.0)	ra  rb  mlx_lm.models.cacher  rR   rU  AttributeError
make_cacheRuntimeError)rT   rd  r  s      r.   _make_mlx_cachezLLMHandler._make_mlx_cache
  sg    
	=$T__55^, 	1133! "Q 		s      A%A	A%A!!A%c                   ]^ ddl m^ ddl}ddlm}m] ddlm} | j                  |ddd      }|d	   }|j                  d
   }^j                  |d         }| j                  |d      }| j                  j                  }| j                  j                  xs |} ||dkD  r|nd|d|cxk  rdk  rn n|nd||dkD  r|nd      }|dk7  }t        |      }|dkD  }|rdnd}d}ddlm}  d}!d}"d}#| j#                  |	|
|ddddddd
      }$|$t%        |$d      rC|$j&                  7^j                  |$j&                  j                         j                               }!t%        |$d      r!|$j                  t)        |$j                        }"t%        |$d      r|$j*                  }#|$j,                  | j.                  k(  rd|v r| j0                  |$_        d|$_        t5        j4                         }%t7        j8                  d| d       ]^fd}&|r| j;                  ||||dd      }'| j                  |'ddd      }(^j                  |(d	   d         })t=        |)      }* || j>                        }+ || j>                        },|}-t=        |-      d
kD  r~tA        |t=        |-      d
z
        }.| j?                  |-d|. d   |+       ^jC                  |+D /cg c]  }/|/j,                   c}/       |-|.d }-^jE                          t=        |-      d
kD  r~|)}0t=        |0      d
kD  r~tA        |t=        |0      d
z
        }.| j?                  |0d|. d   |,       ^jC                  |,D /cg c]  }/|/j,                   c}/       |0|.d }0^jE                          t=        |0      d
kD  r~| j?                  |-d   |+      }1| j?                  |0d   |,      }2^jC                  |1|2       |+g}3|,g}4tG        d
|      D ]0  }5|3jI                   |&|+             |4jI                   |&|,             2 tG        d
|      D ]p  }5 ^jB                  |3|5   D /cg c]  }/|/jJ                  |/jJ                   c}/   ^jB                  |4|5   D /cg c]  }/|/jJ                  |/jJ                   c}/  r |1ddddddf   g|z  }6|2ddddddf   g|z  }7t5        j4                         |%z
  }8||*z   }9|8dkD  r|9|8z  nd}:t7        j8                  d|9 d | d!|* d"|8d#d$|:d%d&| d'|d
z
  |9z   d(       n || j>                        };|}<t=        |<      d
kD  r~tA        |t=        |<      d
z
        }.| j?                  |<d|. d   |;       ^jC                  |;D /cg c]  }/|/j,                   c}/       |<|.d }<^jE                          t=        |<      d
kD  r~| j?                  |<d   |;      }=^jC                  |=       |;g}>tG        d
|      D ]  }5|>jI                   |&|;              tG        d
|      D ]9  }5 ^jB                  |>|5   D /cg c]  }/|/jJ                  |/jJ                   c}/  ; |=ddddddf   g|z  }?t5        j4                         |%z
  }8|8dkD  r||8z  nd}:t7        j8                  d| d)|8d#d$|:d%d&| d*	       tM        |d         }@tG        |      D Acg c]  }AtM        @       }B}AtG        |      D Acg c]  }Ag  }C}Adg|z  }Ddg|z  }Eg }FtG        |      D ]>  }5|r#|5t=        |      k  rFjI                  ||5          (FjI                  d+|5d,z  z          @ t5        j4                         }GtO        |d-| d.| d/d01      }HtG        |      D ]  }ItQ        E      r ntG        |      D ]  }5E|5   r
^jR                  jU                  F|5   Id,z  z          |r7|5   |6|5   |7|5   z
  z  z   }Jn?|5   }JJjW                  d
d      }J|rSt=        B|5         dkD  rB^j                  B|5         }KJdd|Kf   }L^jY                  |LdkD  |L|z  |L|z        }M|M|Jdd|Kf<   |!J|!z   }J|#|"D|5   |#k  rF^j[                  Jddd|"f   ^j                  t        d2      gg      |Jdd|"d
z   df   gd
3      }Jn^Jdd|"|"d
z   f   }N^j]                  |Jj                  t        d2            }J^j[                  |Jddd|"f   |N|Jdd|"d
z   df   gd
3      }JJ^j_                  |Jd4      z
  }O ||O      }P^jC                  |P       |Pja                         }QC|5   jI                  |Q       B|5   jI                  |Q       D|5xx   d
z  cc<   |Q|k(  rdE|5<   |||k7  rQ|k(  rdE|5<   ^j                  Qgg      }R|rP| j?                  R3|5         }S| j?                  |R4|5         }T|Sddddddf   6|5<   |Tddddddf   7|5<   `| j?                  R>|5         }U|Uddddddf   ?|5<    Hjc                  d
       Id5z  dk(  sIdkD  sɉ^jE                           Hje                          t5        j4                         Gz
  }Vtg        d6 CD              }W|dkD  rW|z  nd}XVdkD  rWVz  nd}Y|8Vz   }Zt7        j8                  d7| d8W d9Xd:d;|Vd#d$|Yd%d<|8d#d=|Vd#d>|Zd#d?       g }[tG        |      D ]3  }5| j                  ji                  C|5   d@      }\[jI                  |\       5 [S c c}/w c c}/w c c}/w c c}/w c c}/w c c}/w c c}Aw c c}Aw )Aa  
        Optimized native MLX batch generation for codes phase.

        Strategy: shared prefill + clone cache + interleaved B=1 decode.

        On Apple Silicon, LLM decode is memory-bandwidth-bound. Batching the
        forward pass (B>1) doubles the KV cache reads per step and actually
        *slows down* throughput for 1.7B-class models. Instead, we:

        1. Prefill ONCE with B=1, then clone the KV caches for each item.
           This saves ~50% of prefill time vs sequential generation.
        2. Interleave B=1 forward passes across items within each step.
           Each item gets its own cache, constrained state, and seed.

        This achieves ~1.25x speedup over fully sequential generation while
        maintaining the full ~44 tok/s per-item decode speed.

        Only used for codes generation phase where all prompts are identical.
        Raises on failure so the caller can fall back to sequential mode.
        r   N)r  KVCachemake_samplernpFTro  r  r   r   r   r   r   r   tempr   r   CFG r    rP  FSMStaters  non_audio_code_maskr
  r   r3   z&MLX batch native: prefilling once for z items (shared prompt)c                    g }| D ]w  }        }|j                   Qj                  |j                         |_         j                  |j                        |_        |j                  |_        |j	                  |       y |S )zNDeep-copy a list of KVCache objects so each batch item gets independent state.)r  arrayvaluesoffsetr|   )
cache_listclonedcnew_cr  rd  s       r.   _clone_cache_listz;LLMHandler._run_mlx_batch_native.<locals>._clone_cache_listJ  si    F	66%!#!&&!1EJ#%88AHH#5EL#$88ELe$   Mr0   r_  cacher   zMLX batch native prefill:  tokens (cond=	, uncond=) in r)  s (r   z tok/s) [shared across z items, saved z redundant tokens] tokens in z items]*   iCB MLX zBatch Gen (native, n=r+  toktotalr-  r.  r   axiskeepdims   c              3   2   K   | ]  }t        |        y wr  )rb  ).0ts     r.   	<genexpr>z3LLMHandler._run_mlx_batch_native.<locals>.<genexpr>+  s     ;?a3q6?s   z&MLX batch native generation complete: z items, z total tokens (z.0fz	 avg) in  tok/s) | prefill s + decode s = s totalrz  )5ra  rb  numpyr  r  r  mlx_lm.sample_utilsr  r?   r  r  r   r
  r  r   $acestep.constrained_logits_processorr  r   rL   r  r   r   state	THINK_TAGCODES_GENERATIONr  r4  r
   r   r   rb  rR   r   r   clear_cacher  r|   r  r"  r   allr  seedreshapewhereconcatenatefull	logsumexpr   rs  closer  r  )_rT   rm  rf  r  rY  r   r   r   r   r   r   r   r   r   r   rZ  r  r  r  r  input_ids_npprompt_lengthr  r   r
  r  sampleruse_rep_penaltyrep_penalty_valuse_cfg	cfg_labelprefill_step_sizer  _mlx_non_audio_mask_mlx_eos_id_target_codesrP   prefill_startr  uncond_textuncond_inputsuncond_promptuncond_lengthbase_cond_cachebase_uncond_cachecond_remaining
chunk_sizer  uncond_remainingbase_cond_logitsbase_uncond_logitsitem_cond_cachesitem_uncond_cachesr  item_last_conditem_last_uncondprefill_timetotal_prefill_tokensprefill_tps
base_cache	remainingbase_logitsitem_cachesitem_last_logitsbase_token_idsrH  item_all_token_idsitem_new_tokensitem_codes_countitem_finisheditem_seed_basesdecode_startpbarr7  step_logitstoken_indicesselectedmodifiedeos_vallogprobs	token_arrtoken_id
next_inputrA  rB  
logits_outdecode_timetotal_tokens
avg_tokens
decode_tpsr  rj  r  r  rd  s_                                                                                                @@r.   _run_mlx_batch_nativez LLMHandler._run_mlx_batch_native
  s   L 	B4 ##	 $ 
 k*$**1-,q/* 55+$ 6 
 ))66))66F,  +aS ,u1Bs1B% ,%
 -3 23c/%F2	  	B" !% A A%='A+#$ !B !
 !,,.CDI^IrIrI~&(hh/D/X/X/^/^/`/f/f/h&i#,n=BWBdBdBp!"7"D"DE,n= 5 B B %**h.@.@@!112:2K2K)/89)5 		<ZLH^_`	 :: /'TX ; K !..D%D / M HH];%?%BCM.M 0@O 1$// B $Nn%) !2C4G!4KL
{
 ;D AY/:/Q/:;!/
!<  n%)  -&'!+ !2C8H4IA4MN
 0* =d CK\]*;<*;Q*;<=#3JK#@   &'!+  $~d/C?[!%1A$1GO`!aGG$&89 !00"3!41j) ''(9/(JK"))*;<M*NO * 1j)*:1*=T*=QAS!&&*=TU*<Q*?V*?Q166CU!&&*?VW *
 /q"#qy9:ZGN 21bc19 =>K99;6L#0=#@ AMPQAQ.=WXKKK,-A,B C&y @"3's;s*; <"",^Z\K_<_;``rt +4??;JIi.1$ !2C	NQ4FG
	+: 6t <JO*5*Q*56%jk2	  i.1$ //)D//LKGGK %,K1j)""#4Z#@A *1j)+a.O.QAFF<N!&&.OP * !,ArsAI 67*D99;6L:F:J-,6PQKKK,]O <"3's;s*; <"",W6 l1o.<A*<MN<Mqd>2<MN',Z'89'8!2'893+*, z"ASZ&&uQx0&&rAK'78	 # yy{.i[@UV`Uaab/cjop.)D=! :& # 		q1D7NBC "21"5	^TUEVYijkYlEl8m"mK"21"5K)11!R8 #s+=a+@'AA'E$&HH-?-B$CM*1m+;<H!xx 1 ?2 ?2 H
 5=K= 01 '2"-0C"CK ,1H'*]:&(nn'<K<8HHuV}o%67';?+;(;<6 !"	 '5 '# #.a[1_1L.L"M&(ggk.?.?v&O&(nn'<K<8#';?+;(;<6 !"	 '5 '# 'kD)QQ#H-		"$>>+"))(3"1%,,X6 !#q(# |+'+M!$+0LQY]iQi'+M!$  XXzl3
"&//*DTUVDW/"XK$(OOJFXYZF[O$\M(3ArsAI(>N1%*723	*B$Q'!%;q>!RJ*4QQY*?$Q'W 'Z KKN czQ4!8 o *r 	

 iikL0;?;;2<q.\J.a
3>?\K/
!K/
4ZLnOJs+;9[QTDU V31,s1C;{[^N__cdnorcssz|	
 z"A,,33OA4F\a3bK, # _ ; =" UV, 6 P, O9sB   l1l6 l;
l;
7m 
	m 
 m2m

m

m-	mc                    ddl m} ddl}ddlm} ddlm} | j                  |ddd      }|d	   }|j                  d
   }|j                  |d         }| j                  ||	|
||||||d
      }| j                  |
|      }| j                  j                  }| j                  j                  xs |} ||dkD  r|nd|d|cxk  rdk  rn n|nd||dkD  r|nd      } |dk7  }!t        |      }"|dkD  }#|#rdnd}$d|$ d}%d}&ddlm}' d}(d})d}*d}+|t#        |d      rC|j$                  7|j                  |j$                  j                         j                               }(t#        |d      r!|j                  t'        |j                        })t#        |d      r|j(                  }*|dk(  rL|j*                  |'j,                  k(  r3d|v r/|'j.                  |_        d|_        d}+t3        j4                  d       t7        j6                         },|#r9| j9                  |||||d      }-| j                  |-ddd      }.|j                  |.d	   d         }/t;        |/      }0 || j<                        }1 || j<                        }2|}3t;        |3      d
kD  r~t?        |&t;        |3      d
z
        }4| j=                  |3d|4 d   |1       |jA                  |1D 5cg c]  }5|5j*                   c}5       |3|4d }3|jC                          t;        |3      d
kD  r~|/}6t;        |6      d
kD  r~t?        |&t;        |6      d
z
        }4| j=                  |6d|4 d   |2       |jA                  |2D 5cg c]  }5|5j*                   c}5       |6|4d }6|jC                          t;        |6      d
kD  r~| j=                  |3d   |1      }7| j=                  |6d   |2      }8|jA                  |7|8       |7ddddddf   }9|8ddddddf   }:t7        j6                         |,z
  };||0z   }<|;dkD  r|<|;z  nd}=t3        j4                  d|< d | d!|0 d"|;d#d$|=d%d&       n || j<                        }>|}?t;        |?      d
kD  r~t?        |&t;        |?      d
z
        }4| j=                  |?d|4 d   |>       |jA                  |>D 5cg c]  }5|5j*                   c}5       |?|4d }?|jC                          t;        |?      d
kD  r~| j=                  |?d   |>      }@|jA                  |@       |@ddddddf   }At7        j6                         |,z
  };|;dkD  r||;z  nd}=t3        j4                  d| d'|;d#d$|=d%d&       tE        |d         }Bg }Ct7        j6                         }DtG        ||%d()      }EtI        |      D ])  }F|#r:|9|:z
  z  z   }GnA}GGjK                  d
d      }G|!rMt;        B      dkD  r?|j                  B      }HGdd|Hf   }I|jM                  |IdkD  |I|"z  |I|"z        }J|J|Gdd|Hf<   |~|j*                  }K|K|'j.                  k(  r|(G|(z   }G|*Y|)V|j0                  |*k  rG|jO                  Gddd|)f   |j                  t        d*      gg      |Gdd|)d
z   df   gd
+      }Gn Gdd|)|)d
z   f   }L|jQ                  |Gj                  t        d*            }G|jO                  |Gddd|)f   |L|Gdd|)d
z   df   gd
+      }GnK|'jR                  k(  rnGjU                  |jV                        }M|j                  |Md,      }NtY        jZ                  |N      }OtY        j\                  BgtX        j^                  -      }P ||P|O      }O|j                  |Oj                               }GG|ja                  |Gd.      z
  }Q | |Q      }R|jA                  |R       |Rjc                         }SCje                  |S       Bje                  |S       Ejg                  d
       ||ji                  S       S|k(  r n|||k7  rS|k(  r n|j                  Sgg      }T|#rC| j=                  T1      }7| j=                  |T2      }8|7ddddddf   }9|8ddddddf   }:n!| j=                  T>      }@|@ddddddf   }AFd/z  dk(  sFdkD  s|jC                          , Ejk                          t7        j6                         Dz
  }Ut;        C      }V|UdkD  rVUz  nd}W|;Uz   }Xt3        j4                  d0V d'|Ud#d$|Wd%d1|;d#d2|Ud#d3|Xd#d4       | j                  jm                  Cd5      }Y|YS c c}5w c c}5w c c}5w )6a  
        Optimized native MLX generation using mlx-lm infrastructure.

        Key improvements over the hybrid approach:
        1. Native MLX sampling (temperature, top-k, top-p) via mlx-lm make_sampler
           - Eliminates numpy/PyTorch round-trip for EVERY generated token
        2. Native MLX repetition penalty (no per-step PyTorch conversion)
        3. Chunked prefill for memory-efficient long prompt processing
        4. Periodic memory cleanup (mx.clear_cache) matching mlx-lm patterns
        5. Bridges to PyTorch ONLY for constrained decoding FSM when active

        Raises on failure so the caller can fall back to the legacy hybrid method.
        r   Nr  r  r  FTro  r  r   rs  r  r   r   r  r  r    r  zGen (native)rP  r  r  r
  r   r   r3   zGMLX native: pre-transitioned FSM to CODES_GENERATION (native fast path)r_  r  r   zMLX native prefill: r  r  r  r)  r  r    tok/s)r  r  r  r   r  r  rF   r  r  z MLX native generation complete: r  r  r  r  rz  )7ra  rb  r  r  r  r  r  r?   r  r  r   r   r
  r  r   r  r  rL   r  r   r   r  r  r  r  r
   r   r4  r   rb  rR   r   r   r  r"  r   r  r  r  r  r  	COMPLETEDastyperE   rD   
from_numpytensorlongr  r   r|   rs  r  r  r  )ZrT   rm  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rd  r  r  r  r  r  r  r  rP   r   r
  r  r  r  r  r  r  	tqdm_descr  r  r  r  r  _use_native_codes_pathr  r  r  r  r  
cond_cacheuncond_cacher  r  r  r  rA  rB  	last_condlast_uncondr  r  r  r  r  r  last_logitsall_token_ids
new_tokensr  r  r7  r  r  r  r  	_cp_stater  step_logits_f32	np_logitst_logitst_idsr  r  r  r  r  num_generatedr  r  r  sZ                                                                                             r.   _run_mlx_single_nativez!LLMHandler._run_mlx_single_native=  ss
   F 	94 ##	 $ 
 k*$**1-,q/* !% A A%='A+'/#%'- !B !
 55+- 6 
 ))66))66F,  +aS ,u1Bs1B% ,%
 -3 23c/%F2	9+\2	 
 	B"!& ,,.CDI^IrIrI~&(hh/D/X/X/^/^/`/f/f/h&i#,n=BWBdBdBp!"7"D"DE,n= 5 B B
  7*/D/J/JhN`N`/`!112:2K2K)/89)5-1*KK ij 		::! /!1 ; K !..#	 / M HH];%?%BCM.M +4??;J,T__=L $Nn%) !2C4G!4KL
{
 ;D AT*5*Q*56!/
!<  n%)  -&'!+ !2C8H4IA4MN
 0* =d C<X,7,Q,78#3JK#@   &'!+ //.*>j/QK OO,<T,B,OWMGGK/#ArsAI.I'23	2K99;6L#0=#@ AMPQAQ.=WXKKK&';&< =&y @"3's;s*;7D &doo6E Ii.1$ !2C	NQ4FG
	+: 6t <EJ%0%Q%01%jk2	  i.1$ 4FJGGJ$QQY/K99;6L:F:J-,6PQKKK&}o 6"3's;s*;7D \!_-
yy{.yuE.)D)I[9P,QQ)%--a4K 3}#5#9 " 7&q-'7888qL..
 19A},- %0177	 9 99 +6&14G&G$0[5L0<<}L*,.. +A||O < "5=/): ; +A{Q/?,? @: %&	 +9 +'K '2![q5P2P&QG*,''+2C2CU6]*SK*,.. +A||O < ' +A{Q/?,? @: %&	 +9 +'K ("4"44 '2&8&8&DO "t DI$//	:H!LL-

KE4UHEH"$((8>>+;"<K #R\\+\%MMH)IGGI ~~'Hh'  *KKN %0%228< <''LL,HXYeMe
 H:,/J"ooj
oK $
, O'23	2	+ArsAI6!__Zu_E
(BC3 czQ4!8 A *D 	

 iikL0J4?!O][0
!K/
.}o[UXHY Z31,s1C;{[^N__cdnorcssz|	
 ((//
PU/VW 6 8: 1s   ee)e#c                 $   	 | j                  |||||||||	|
|||||||||      S # t        $ r9}t        j                  dt	        |      j
                   d| d       Y d}~nd}~ww xY wddlm} ddl}| j                  |ddd	
      }|d   }|j                  d   }|j                  |      }| j                  ||	|
||||||d
      }| j                  |
|      }| j                  j                  }| j                  j                  xs |}|dkD  }|rdnd} d|  d}!t!        j                          }"|r| j#                  |||||d      }#| j                  |#ddd	
      }$|j                  |$d         }%|%j                  d   }&| j%                         }'| j%                         }(| j'                  ||'      })| j'                  |%|(      }*|j)                  |)|*       |)ddddddf   }+|*ddddddf   },t!        j                          |"z
  }-||&z   }.|-dkD  r|.|-z  nd}/t        j*                  d|. d| d|& d|-dd|/dd       n| j%                         }0| j'                  ||0      }1|j)                  |1       |1ddddddf   }2t!        j                          |"z
  }-|-dkD  r||-z  nd}/t        j*                  d| d|-dd|/dd       t-        |d         }3g }4t!        j                          }5t/        ||!d !      }6t1        |      D ]  }7|r,|+|,z
  z  z   }8n2}8|8j3                  dd      }8|8j5                  |j6                        }9|j                  |9d	"      }:t9        j:                  |:      };t9        j<                  |3gt8        j>                  #      }<|	 ||<|;      };|dk7  rdd$l m!}=  |=|%      }> |>|<|;      };| jE                  |;|      };| jG                  |;|      };| jI                  |;|      }?|?jK                         }@|4jM                  |@       |3jM                  |@       |6jO                  d       ||jQ                  @       @|k(  r n|||k7  r@|k(  r n|j                  @gg      }A|rV| j'                  A'      })| j'                  |A(      }*|j)                  |)|*       |)ddddddf   }+|*ddddddf   },| j'                  A0      }1|j)                  |1       |1ddddddf   }2 |6jS                          t!        j                          |5z
  }BtU        |4      }C|BdkD  rCBz  nd}D|-Bz   }Et        j*                  d&C d|Bdd|Ddd'|-dd(|Bdd)|Edd*       | j                  jW                  |4d+      }F|FS ),a  
        MLX-accelerated single-item generation.

        Tries optimized native MLX generation first (using mlx-lm infrastructure
        for sampling, repetition penalty, and chunked prefill). Falls back to
        hybrid MLX/PyTorch approach if native generation fails.
        r  zNative MLX generation failed (r#  z), falling back to hybrid modeNr   r  FTro  r  r   rs  r  r   r  r    r  
Generationr_  r  r   zMLX prefill: r  r  r  r)  r  r   r  r  r  r  r  r  )r   r   zMLX generation complete: r  r  r  r  rz  ),r  r^   r
   rj   r  r%  ra  rb  r  r?   r  r  r   r   r
  r  r4  r   r  rR   r   r   r"  r   r  r  r	  rE   rD   r
  r  r  &transformers.generation.logits_processr   r   r   r  r   r|   rs  r  r  rb  r  )GrT   rm  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   _native_errrd  r  r  r  r  r  rP   r   r
  r  r  r  r  r  r  r  r  r  r  r  rA  rB  r  r  r  r  r  r  r  r  r  r  r  r  r7  r  r  r  r  r  r   rep_proct_tokenr  r  r  r  r  r  r  sG                                                                          r.   _run_mlx_singlezLLMHandler._run_mlx_single  s   <	..!1'# /#5)A+E /+"3')+!1!' /  *  	NN0k1B1K1K0LB{m \. / 	 	 ##	 $ 
 k*$**1-,' !% A A%='A+'/#%'- !B !
 55+- 6 
 ))66))66F,c/%F2	9+Z0	 		::! /!1 ; K !..#	 / M HH];%?@M)//2M --/J//1L //&
/CK OOMONMGGK/#ArsAI.I'23	2K99;6L#0=#@ AMPQAQ.=WXKKK 45 6&y @"3's;s*;7D ((*Eu=JGGJ$QQY/K99;6L:F:J-,6PQKKK /"3's;s*;7D \!_-
yy{.yuE.)D)I[9P,QQ)%--a4K
 *00<Ot<I''	2HLL-

CE %00A "S(c;DVW#E84 //%@H//%@H ))(K@G||~Hh'  *KKN %0%228< <''LL,HXYeMe H:,/J"ooj
oK $
, O]3'23	2	+ArsAI6!__Zu_E

#(BC3y *| 	

 iikL0J4?!O][0
!K/
'k+cAR S31,s1C;{[^N__cdnorcssz|	
 ((//
PU/Vs   #& 	A(/A##A(c                 \   ddl m} | j                  |      \  }}|r"t        |      }t        t	        |            dk(  }|dk(  xr' |xr# |dkD  xr t        | d      xr | j                  du}|r=	 t        j                  d| d       | j                  |d   |||||||||	|
||||      S t        j                  d| d       g }t        |      D ]h  \  }}|r,|t        |      k  r|j                  j!                  ||          | j#                  |||||||||	|
ddddd||||      }|j%                  |       j |S |d   }| j#                  |||||||||	|
|||||||||      S # t        $ r9}t        j                  d	t        |      j                   d
| d       Y d}~d}~ww xY w)a9  
        Unified MLX generation function supporting both single and batch modes.

        For batch mode in codes generation phase, uses optimized batch native path
        that shares prefill across all items (saving ~50% prefill time).
        Falls back to sequential processing if batch native fails.
        r   Nr   r   rR   z9MLX batch: using optimized batch native path (batch_size=z, shared prefill))rm  rf  r  rY  r   r   r   r   r   r   r   r   r   r   rZ  zMLX batch native failed (r#  z"), falling back to sequential modez-MLX batch: using sequential mode (batch_size=r+  FTr  )ra  rb  r#  rb  setrL   rR   r
   r   r  r^   rj   r  r%  r  r  r  r#  r|   ) rT   r  r  rY  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rZ  rd  re  r   rf  all_prompts_identicalcan_use_batch_nativer   rj  r  rm  r  s                                    r.   r  zLLMHandler._run_mlxp  sa   < 	 +/*E*EFW*X'x23J %(,A(B$Cq$H! G+ 0)0N0 D,/0 OO4/ ! $KK''1l2CE  55)>q)A#-$/"+(7##+=1I3M(7 '%!)# 6  0 KKG
|STUVL'01F'G##QU^IINN58,"22%5 +'$3'9-E/I$3"&&+ $!%"&%5#!%' 3 * ##K05 (H6   13##-#+1%='A+'/#%'-' $ 
 	
Q ! NN3DG4D4D3ERs K: ; s   1;E) )	F+2/F&&F+c              #   	  K   | j                   sd y| j                  dv rd y| j                  }|d y	 t        |j	                               j
                  j                  }t        | j
                        j                  d      d   }||k(  rd yt        j                  d| j
                          t        j                         }t        |d      r4|j                  | j
                        j                  | j                         t        j                         |z
  }t        j                  d| j
                   d|d	d
       	 d t        j                  d       t        j                         }t        |d      r|j                  d       t         j"                  j%                         rt         j"                  j'                          nt        t         d      r=t         j(                  j%                         rt         j(                  j'                          nt        t         j*                  d      rpt         j*                  j,                  j%                         rHt        t         d      r8t        t         j,                  d      rt         j,                  j'                          t        j                         |z
  }t        j                  d|d	d
       y# t        $ r d}Y w xY w# t        j                  d       t        j                         }t        |d      r|j                  d       t         j"                  j%                         rt         j"                  j'                          nt        t         d      r=t         j(                  j%                         rt         j(                  j'                          nt        t         j*                  d      rpt         j*                  j,                  j%                         rHt        t         d      r8t        t         j,                  d      rt         j,                  j'                          t        j                         |z
  }t        j                  d|d	d
       w xY ww)z
        Context manager to load a model to GPU and offload it back to CPU after use.
        Only used for PyTorch backend when offload_to_cpu is True.
        N)rX   r/  rN  r   zLoading LLM to r   zLoaded LLM to z in z.4frn  zOffloading LLM to CPUr8   r]   rZ   r\   zOffloaded LLM to CPU in )rG   rA   r>   rt  rx  rC   r  StopIterationr   r+   r
   r   r4  rL   r   rF   rD   rb   rc   r\   r]   rd   rZ   )rT   r  rV  target_devicer<  r  offload_times          r.   r|  zLLMHandler._load_model_context  s<     "" .=
	"!%"2"2"45<<AAN DKK(..s3A6]* 	odkk]34YY[
5$HHT[[!$$TZZ0IIK*,	nT[[Mi_AFG	H KK/1Jud#zz&&(

&&(&599+A+A+C		%%'/ENN4F4F4S4S4UZabginZot{  }B  }F  }F  HU  uV		%%'99;3LKK2<2DAFG?  	"!N	"& KK/1Jud#zz&&(

&&(&599+A+A+C		%%'/ENN4F4F4S4S4UZabginZot{  }B  }F  }F  HU  uV		%%'99;3LKK2<2DAFGsC   8R -K (C%R K0 FR K-)R ,K--R 0FQ==R c                    | j                   dk(  r| j                  S | j                   dk(  r| j                  qt        j                  d       | j                  j
                  }|j                  j                  }ddl} |j                         }t        j                  |d| j                        | _         |j                         |z
  }t        j                  d|d	d
       | j                  r;| j                  j                          t        j                  d       | j                  S t        |j                  j                               j                   }| j                  j#                  |      | _        | j                  j                          t        j                  d|        | j                  S | j                   dk(  r| j                  xt        j                  d       t%        | dd      }|t'        d      ddl} |j                         }t        j                  |d| j                        | _         |j                         |z
  }t        j                  d|d	d
       | j                  r;| j                  j                          t        j                  d       | j                  S t)        t*        j,                  d      r*t*        j,                  j.                  j1                         rdnd}| j                  j#                  |      | _        | j                  j                          t        j                  d|        | j                  S t'        d| j                          )aw  
        Get HuggingFace model for perplexity scoring.

        For vllm backend, loads HuggingFace model from disk (weights are cached by transformers).
        For pt backend, returns the existing model.
        For mlx backend, loads HuggingFace model from disk (MLX model can't be used for torch scoring).

        Returns:
            HuggingFace model instance
        r   rX   Nz7Loading HuggingFace model for scoring (from checkpoint)r   T)r   torch_dtypezHuggingFace model loaded in r)  rn  z?HuggingFace model for scoring kept on CPU (offload_to_cpu=True)z'HuggingFace model for scoring ready on r/  zGLoading HuggingFace model for scoring (MLX backend, need PyTorch model)rS   zEMLX model path not stored. Cannot load HuggingFace model for scoring.rZ   r8   zUnknown backend: )rA   r>   rQ   r
   r   model_runnerr~  r  r4  r   r   rF   rG   r   rt  rx  rC   r   r%   r  rL   rD   rd   rZ   rc   )rT   r.  r   r4  r<  r  rC   s          r.   get_hf_model_for_scoringz#LLMHandler.get_hf_model_for_scoring4  s    t#88O' ))1UV  $xx44)0066
 &TYY[
-A-Q-Q&* $

.*
 &DIIK*4	:9S/KL &&..335KK ab --- ","4"4"?"?"ABIIF151K1K1N1Nv1VD...335KK"I& RS---&))1ef %T+<dC
%$%lmm&TYY[
-A-Q-Q&* $

.*
 &DIIK*4	:9S/KL &&..335KK ab --- '.enne&DI[I[IhIhIjUpuF151K1K1N1Nv1VD...335KK"I& RS--- 01A1A0BCDDr0   r  )rV   N)Nr   g?r   )FNN)F)rX   r'  FN)TFNNNNFTFFr   r    r    r    N)TFNNFTFFr   r    r    r    N)333333?r   r   NNr   TFNNTTTNNN)r    Fr   r   )Fr   )g333333?NNr   TF)FFr   )FNr0  NNr   TF)Nr0  NNr   TF)NTFF)Lr%  
__module____qualname____doc__STOP_REASONING_TAGrH   rI   rJ   rO   r   r   rU   re   r_   rw   r   r   r   r   r6  r   r   r   r   r   r   r   r   r   r   r   rD   r  r   r   r  r  r  r   r  r   r#  rF   rJ  r:  rl  r  r  r  r  r  r   r   r  r  r  r  r  r  r!  r  r   r  r  r  staticmethodr7  r8  r  r  r  r#  r  r	   r|  r/  r  r0   r.   r2   r2   .   s5   2# ::>>*5TA$ $6"HY9S 9T#Y 0S 0e 0di 0  }B 0  MR  SX  Z^  S^  M_ 0l '+	:!%: : sm	:
 
:xjs jt j %  DW  $ 04-18*"&8* %)8* "%	8*
  S(3-%7 898*  8* 8* 8* 8* 8* 8* 'uo8* $E?8* 
4	58*B   	
    
,mc m3 m5sCS m"%,, x} QVQ]Q] %,, x SXS_S_ 0U\\ 0 0%,, 0u|| 3 V^_bVc hl ERtIu E  @E  @L  @L E || 38n	
 "#  
01c49n8M 1RWX\]`XacgXgRh 1 $'+UmUm Um 	Um
 Um Um $Um 
sDy	Umn2f# 2fd 2fWZ 2fz *.+004-1+/<@"' "# %%)/mA d3i0mA mA 	mA
 mA }mA mA "mA #'mA %)mA 'uomA $E?mA "%mA  S(3-%7 89mA  mA  !mA" #mA$ %mA& 'mA( )mA* +mA, -mA. S	"/mA0 
sDI~	1mA^^^ ^ 	^
 ^ }^ ^ "^ #'^ %)^ "%^  S(3-%7 89^  ^ ^ ^  !^" #^$ %^& '^( )^* 
+^R *.+0+/<@"' "# %%)+b
 d3i0b
 b
 	b

 b
 }b
 b
 "b
 #'b
 %)b
 "%b
  S(3-%7 89b
  b
 b
 b
  !b
" #b
$ %b
& 'b
( )b
* S	"+b
, 
sDI~	-b
H8Dhsm9K4L+M RV 0S#X 03 0D ".#!%$')-+0+/<@" $!%$(%))jj j 	j
 j j j }j j "j #'j %)j "%j  S(3-%7 89j j  !j" #j$ SM%j& S	"'j* 
c3h+jX.
c .
3 .
Y] .
y| .
  X[ .
  ru .
`<s <C <SV <lp <  LO <  fi <B $).	.
.
 !.
 	.

 
.
f !#!%$')-+0b$b$ b$ }	b$
 b$ "b$ #'b$ %)b$ 
tCH~s"	#b$H$#s $#s $#R ##(.4
4
 4
 !	4

 4
 
4
r #(,!#!%$')-+0x$x$ x$ !	x$
 x$ }x$ x$ "x$ #'x$ %)x$ 
tCH~s"	#x$| $).2
2
 2
 !	2

 2
 
2
p 37!#!%$')-+0P$P$ P$  S#X/	P$
 P$ }P$ P$ "P$ #'P$ %)P$ 
tCH~s"	#P$j )-)-+0"'bb d38n%b #'	b
 %)b  b 
sCxb^ OS[<<[ !.[ 	[
 [ }[ [ "[ [ <([  ((JK[ 
[R OSww 'u||4w 	w
 w w }w w "w w <(w  ((JKw 
wrv%3 v%5c3h9L3M v%x t  V;# V;%c	2B V;p> &*!gg g 	g
 g g }g g "g #'g %)g "%g g g g  S	"!g" 
c#gRAA A 	A
 A }A A "A #'A %)A "%A  S(3-%7 89A  A A A  !A" #A$ %A& 'A( )A* 
+AF
nn n 	n
 n }n n "n #'n %)n "%n  S(3-%7 89n  n n n  !n" #n$ %n& 'n( )n* 
+nr *.+0+/<@"' "# %%)+D
 d3i0D
 D
 	D

 D
 }D
 D
 "D
 #'D
 %)D
 "%D
  S(3-%7 89D
  D
 D
 D
  !D
" #D
$ %D
& 'D
( )D
* S	"+D
, 
sDI~	-D
T 7H 7HrVEr0   r2   ).r3  rH   r#   r   r4  r  r(   typingr   r   r   r   r   r   
contextlibr	   r  rD   logurur
   r   transformersr   r   !transformers.generation.streamersr   r  r   r   r  r   acestep.constantsr   r   r   r   r   r   acestep.gpu_configr   r   r   r   r9  r/   r2   r  r0   r.   <module>r=     sq    
 
     : : %     < : T u  u s s
  	
\=E \=Er0   