
    xi                        d Z ddlmZmZ ddlmZmZmZmZm	Z	m
Z
mZ ddlmZ ddlmZ ddlmZ ddlZddlZddlmZmZmZmZmZmZmZmZmZmZ d	Z G d
 de      Z  G d de      Z!y)aK  
Constrained Logits Processor for ACE-Step Language Model

This module implements a finite state machine (FSM) based logits processor that constrains
the language model's output to follow specific formats and value ranges during music generation.

Key Features:
- Enforces structured metadata generation (BPM, duration, keyscale, etc.)
- Validates numeric ranges (BPM: 30-300, Duration: 10-600s)
- Ensures proper formatting for musical metadata
- Prevents generation of invalid tokens or formats
- Supports constrained audio code generation (0-63999)

The FSM guides the model through different states to ensure outputs conform to expected
schema requirements without post-processing corrections.

Usage:
    processor = ConstrainedLogitsProcessor(tokenizer, mode="metadata")
    outputs = model.generate(inputs, logits_processor=[processor])
    )Enumauto)OptionalDictAnyTupleListCallableSet)logger)AutoTokenizer)LogitsProcessorN)
VALID_LANGUAGESKEYSCALE_NOTESKEYSCALE_ACCIDENTALSKEYSCALE_MODESVALID_KEYSCALESBPM_MINBPM_MAXDURATION_MINDURATION_MAXVALID_TIME_SIGNATURESi  c                   `   e Zd ZdZ e       Z e       Z e       Z e       Z e       Z	 e       Z
 e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Z e       Zy)FSMStatez3Finite State Machine states for metadata generationN)__name__
__module____qualname____doc__r   	THINK_TAGNEWLINE_AFTER_THINKBPM_NAME	BPM_VALUENEWLINE_AFTER_BPMCAPTION_NAMECAPTION_VALUEDURATION_NAMEDURATION_VALUENEWLINE_AFTER_DURATIONGENRES_NAMEGENRES_VALUENEWLINE_AFTER_GENRESKEYSCALE_NAMEKEYSCALE_VALUENEWLINE_AFTER_KEYSCALELANGUAGE_NAMELANGUAGE_VALUETIMESIG_NAMETIMESIG_VALUENEWLINE_AFTER_TIMESIGTHINK_END_TAGCODES_GENERATION	COMPLETED     C/mnt/workspace/ACE-Step-1.5/acestep/constrained_logits_processor.pyr   r   5   s    =I&vHI6LFMFMVN!V&K6L6FMVN!VFMVN6LFM FFMvIr8   r   c                      e Zd ZdZ	 	 	 	 	 dVdedededee   dedee   fd	Z	d
edee
   fdZd ZdefdZdefdZdefdZededefd       ZdefdZdefdZdWdeeeee   f      fdZd Zd Zdedee   fdZd Zdej6                  d ee   ddfd!Zdeeed"f   ee   f   fd#Z 	 	 dXd$ee   d%ed&edeeed"f   ee   f   fd'Z!deeed"f   ee   f   fd(Z"d) Z#d* Z$d+ Z%defd,Z&dYd-ed.ed/e'd0efd1Z(d2 Z)d3 Z*d.edee   fd4Z+d5edefd6Z,d7ed.edee   fd8Z-dee   fd9Z.d: Z/d;ee0   fd<Z1defd=Z2d>edee   fd?Z3d@edAedee   fdBZ4dCeeed"f   ee   f   dee   fdDZ5dEej6                  d@edAedefdFZ6dEej6                  defdGZ7dee   fdHZ8defdIZ9dee   fdJZ:dee   fdKZ;dLejx                  dejz                  dejz                  fdMZ>dLejx                  defdNZ?dejz                  dejz                  fdOZ@dPedeee      fdQZAdLejx                  dejz                  dejz                  fdRZBdS ZCdTefdUZDy)Z"MetadataConstrainedLogitsProcessoru8  
    FSM-driven LogitsProcessor that constrains generation to produce valid metadata.
    
    This processor enforces the following format:
    <think>
    bpm: [30-300]
    caption: [text without code blocks, ends with period + newline]
    duration: [10-600]
    keyscale: [A-G][#/♭]? [major/minor]
    language: [en/zh/ja/ko/es/fr/de/uk/ru/...]
    timesignature: [2/3/4/6]
    </think>
    
    It uses token masking (setting invalid token logits to -inf) to enforce constraints.
    For numeric fields, it uses early-blocking to prevent out-of-range values.
    For field transitions (e.g., end of numeric value), it compares P(newline) vs P(digit).
    For caption field, it blocks code blocks and newlines, and only transitions when
    the previous token was a period and newline has the highest probability.
    N	tokenizerenableddebuggenres_vocab_pathskip_genresmax_durationc                    || _         || _        || _        || _        ||nt        | _        d| _        d| _        d| _        dddddddd| _	        d| _
        d| _        d| _        d| _        d| _        d| _        d| _        t"        j$                  | _        d| _        d| _        g | _        d| _        d| _        d| _        d| _        g | _        d| _        | j;                          |xs^ t<        j>                  jA                  t<        j>                  jC                  t<        j>                  jE                  tF                    d      | _$        g | _%        d| _&        i | _'        i | _(        g | _)        i | _*        | jW                          tX        tZ        d	t\        | j
                  d	d
t^        id| _0        tc        | j`                  d   d   | j`                  d   d   dz         D cg c]  }te        |       c}| _3        tc        | j`                  d   d   | j`                  d   d   dz         D cg c]  }te        |       c}| _4        | j`                  d   d
   D cg c]  }te        |       c}| _5        | jm                         | _7        | jq                  | jf                  dd      | _9        | jq                  | jh                  dd      | _:        | jq                  | jj                  dd      | _;        | jy                         | _=        | j}                          t"        j$                  dt"        j~                  dt"        j                  dt"        j                  dt"        j                  dt"        j                  dt"        j                  dt"        j                  dt"        j                  dt"        j                  di
| _H        | j                          yc c}w c c}w c c}w ) aT  
        Initialize the constrained logits processor.
        
        This processor should be initialized once when loading the LLM and reused
        for all generations.
        Args:
            tokenizer: The tokenizer to use for encoding/decoding
            enabled: Whether to enable constrained decoding
            debug: Whether to print debug information
            genres_vocab_path: Path to genres vocabulary file
            skip_genres: Whether to skip genres field generation
            max_duration: Maximum duration in seconds (default: DURATION_MAX from constants)
        NFbpmcaptiondurationkeyscalelanguagetimesignaturegenresr   cot zgenres_vocab.txtg        )minmaxvalid_values)rD   rF   rI   rD   rM   rN      rF   rI   zbpm:bpm: context_prefix_for_matchingcontext_prefix_for_tokenization	duration:
duration: ztimesignature:timesignature: z<think>
zcaption:zgenres:	keyscale:	language:</think>)Jr<   r=   r>   r@   r   rA   skip_captionskip_languagerE   user_provided_metadatametadata_temperaturecodes_temperaturetarget_durationtarget_codescodes_countstop_at_reasoninggeneration_phaser   r   stateposition_in_stateaccumulated_valueaccumulated_token_idscaption_after_newlinecaption_token_countcaption_endingpending_field_nameuser_field_token_queuecurrent_user_field_precompute_tokensospathjoindirnameabspath__file__r?   genres_vocabgenres_vocab_mtimegenres_triecaption_genres_triecaption_matched_genres_char_to_tokens_precompute_char_token_mappingr   r   r   r   field_specsrangestrvalid_bpm_valuesvalid_duration_valuesvalid_timesig_values_build_keyscale_prefix_treekeyscale_prefix_tree_build_numeric_prefix_treebpm_prefix_treeduration_prefix_treetimesig_prefix_tree_build_language_prefix_treelanguage_prefix_tree_load_genres_vocabr    r!   r$   r&   r)   r,   r/   r1   r4   fixed_strings_build_state_transitions)selfr<   r=   r>   r?   r@   rA   vs           r9   __init__z+MetadataConstrainedLogitsProcessor.__init__f   s   , #
& -9,DL,!"&* !A
# 6:!26 15+/ ! (- &+ ''
!"!#02" &+"#$ #"$ 24#15 	! "3 "
bggllGGOOBGGOOH568J7
 (*),!#)+ 13#/1 	++-
 #73 ,T5F5FG,.CD
 27t7G7G7Nu7UW[WgWghmWnotWuxyWy1z {1zAQ1z {6;D<L<LZ<XY^<_aeaqaqr|a}  D  bE  HI  bI  7J  &K  7Jc!f  7J  &K"595E5Eo5VWe5f$g5fSV5f$g! %)$D$D$F!
  $>>!!(.,3  ?  

 %)$C$C&&(3,8 %D %
!
 $(#B#B%%(8,= $C $
  %)$D$D$F!! 	(($v!!:""K  )""K""K!!#3""J
 	%%'c !| &K$gs   $O/OOcurrent_fieldreturnc                     g d}t         j                  t         j                  t         j                  t         j                  t         j
                  t         j                  t         j                  d}	 |j                  |      }t        |dz   t        |            D ]C  }||   }|dk(  r| j                  r|dk(  r| j                  r,|dk(  r| j                  r>||   c S  t         j                  S # t        $ r t         j                  cY S w xY w)a  
        Get the next field state. Always returns the next field's NAME state,
        even if the field is user-provided (we still need to generate the field name).
        
        Args:
            current_field: Current field name ("bpm", "caption", "duration", "genres", "keyscale", "language", "timesignature")
            
        Returns:
            Next FSMState (NAME state of next field), or THINK_END_TAG if no more fields
        )rD   rE   rF   rJ   rG   rH   rI   rP   rJ   rE   rH   )r   r!   r$   r&   r)   r,   r/   r1   index
ValueErrorr4   r   lenr@   r\   r]   )r   r   field_orderfield_to_statecurrent_idxifields          r9   _get_next_field_statez8MetadataConstrainedLogitsProcessor._get_next_field_state  s     h$$,, ..** .. ..%22
	*%++M:K
 {QK(89ANE  T%5%5	!d&7&7
"t'9'9 "%(( : %%%'  	*)))	*s   2C1 1DDc                 
   t         j                  t         j                  t         j                  t         j                  t         j                  t         j
                  t         j
                  t         j                  i| _        t         j                  | j                  t         j                  <   | j                  d      | j                  t         j                  <   | j                  sWt         j                  | j                  t         j                  <   | j                  d      | j                  t         j                  <   t         j                  | j                  t         j                  <   | j                  d      | j                  t         j                  <   | j                  sWt         j                   | j                  t         j"                  <   | j                  d      | j                  t         j                   <   t         j$                  | j                  t         j&                  <   | j                  d      | j                  t         j$                  <   | j(                  sWt         j*                  | j                  t         j,                  <   | j                  d      | j                  t         j*                  <   t         j.                  | j                  t         j0                  <   t         j                  | j                  t         j.                  <   y)z;Build state transition map based on user-provided metadata.rD   rE   rF   rJ   rG   rH   N)r   r   r    r!   r4   r5   r6   
next_stater"   r   r\   r%   r$   r'   r&   r@   r*   r)   r-   r,   r]   r0   r/   r2   r1   r   s    r9   r   z;MetadataConstrainedLogitsProcessor._build_state_transitions5  s     < <(((*;*;""H$=$=%%x'9'9	
 .6-?-?))*.2.H.H.O**+   5=5K5KDOOH1126:6P6PQZ6[DOOH223 3;2I2I../373M3Mj3Y//0 4<4I4IDOOH001595O5OPX5YDOOH112 3;2I2I../373M3Mj3Y//0 !!6>6M6MDOOH2237;7Q7QR\7]DOOH334 2:1G1G--.2:2H2H../r8   skipc                 2    || _         | j                          y)zDSet whether to skip genres generation and rebuild state transitions.N)r@   r   r   r   s     r9   set_skip_genresz2MetadataConstrainedLogitsProcessor.set_skip_genres`  s    %%'r8   c                 2    || _         | j                          y)zESet whether to skip caption generation and rebuild state transitions.N)r\   r   r   s     r9   set_skip_captionz3MetadataConstrainedLogitsProcessor.set_skip_captione  s     %%'r8   c                 2    || _         | j                          y)zFSet whether to skip language generation and rebuild state transitions.N)r]   r   r   s     r9   set_skip_languagez4MetadataConstrainedLogitsProcessor.set_skip_languagej  s    !%%'r8   rE   c                     | s| S | j                  d      }g }|D ]&  }|j                         }|s|j                  |       ( dj                  |      S )a   
        Post-process caption to remove YAML multi-line formatting.
        Converts YAML-style multi-line text (with newlines and leading spaces) 
        to a single-line string.
        
        Example:
            Input:  "An emotional ballad.\n  The track opens with piano.\n  More text."
            Output: "An emotional ballad. The track opens with piano. More text."
        
        Args:
            caption: Raw caption text with possible YAML formatting
            
        Returns:
            Clean single-line caption
        rX    )splitstripappendrs   )rE   linescleaned_lineslinestrippeds        r9   postprocess_captionz6MetadataConstrainedLogitsProcessor.postprocess_captiono  sZ    " N d# Dzz|H$$X.  xx&&r8   stopc                     || _         y)a  
        Set whether to stop generation after </think> tag.
        
        Args:
            stop: If True, generation will stop immediately after </think> tag is generated.
                  If False, generation continues to codes generation phase.
        N)rd   )r   r   s     r9   set_stop_at_reasoningz8MetadataConstrainedLogitsProcessor.set_stop_at_reasoning  s     "&r8   phasec                 8    |dvrt        d|d      || _        y)u	  
        Set the generation phase.
        
        Args:
            phase: "cot" for CoT metadata generation, "codes" for audio codes generation,
                   or "understand" for audio understanding (codes → metadata + lyrics).
                   When phase is "codes" and the input prompt already contains </think>,
                   the FSM will skip metadata generation and go directly to codes generation.
                   When phase is "understand", generate CoT metadata then free-form lyrics.
        )rK   codes
understandzInvalid generation phase: z). Must be 'cot', 'codes', or 'understand'N)r   re   )r   r   s     r9   set_generation_phasez7MetadataConstrainedLogitsProcessor.set_generation_phase  s+     669%Bklmm %r8   metadatac                 j   |i }dD ](  }||v r||   | j                   |<   d| j                   |<   * | j                          | j                  r`| j                   j                         D cg c]
  \  }}|	| }}}|rt	        j                  d|        yt	        j                  d       yyc c}}w )a  
        Set user-provided metadata fields. Fields that are provided will be used directly
        instead of generating. Fields that are None will be generated.
        
        Args:
            metadata: Dictionary with optional fields:
                - "bpm": Optional[str] - e.g., "120"
                - "caption": Optional[str] - e.g., "A melodic piano piece..."
                - "duration": Optional[str] - e.g., "234"
                - "keyscale": Optional[str] - e.g., "G major"
                - "language": Optional[str] - e.g., "en"
                - "timesignature": Optional[str] - e.g., "4"
                - "genres": Optional[str] - e.g., "Pop Rock"
                If None, clears all user-provided metadata.
        NrC   zUser provided metadata fields: z7No user-provided metadata, all fields will be generated)r^   r   r>   itemsr   )r   r   r   kr   provided_fieldss         r9   set_user_metadataz4MetadataConstrainedLogitsProcessor.set_user_metadata  s      H gE 5=e_++E259++E2	 g 	%%'::-1-H-H-N-N-Pb-PTQTUTaq-POb>>OPQVW bs   +
B/6B/c                 D   i | _         t        d      D ]=  }| j                  j                  t	        |      d      }|s,|d   | j                   |<   ? | j                  j                  dd      }|r|d   nd| _        i | _        t        D ]4  }| j                  j                  |d      }|s#|d   | j                  |<   6 g | _        dD ]@  }| j                  j                  |d      }|s#| j                  j                  |d          B g | _
        dD ]@  }| j                  j                  |d      }|s#| j                  j                  |d          B | j                  j                  d	d      }|r|d   nd| _        g | _        g | _        d
D ]r  }| j                  j                  |d      }|s#|j                         dk(  s7| j                  j                  |d          | j                  j                  |d          t t        | j                        | _        | j                  j                  dd      }	|	r|	d   nd| _        | j                  j$                  | _        | j                  j                  dd      }
|
r|
d   nd| _        | j                  j                  dd      }|r|d   nd| _        t*        | _        t/               | _        | j3                          d| _        d| _        | j9                          t;        j<                         | _        y)z3Pre-compute commonly used token IDs for efficiency.
   Fadd_special_tokensrX   N)#u   ♯)bu   ♭r   )mMr   ,.`) digit_tokensr   r<   encoder   newline_tokennote_tokensr   sharp_tokensr   flat_tokensspace_tokenmajor_start_tokensminor_start_tokenslowerr   
vocab_sizecomma_tokeneos_token_idperiod_tokenbacktick_tokenr   valid_languagessetaudio_code_token_ids_precompute_audio_code_tokensaudio_code_masknon_audio_code_mask_build_audio_code_maskr   copyvalid_keyscales)r   dtokensnewline_tokensnotesfspace_tokensprefixcomma_tokensperiod_tokensbacktick_tokenss               r9   rp   z5MetadataConstrainedLogitsProcessor._precompute_tokens  s    rA^^**3q6e*LF'-bz!!!$  ..t.N3A^B/t "D^^**4E*JF)/  & # A^^**1*GF!!((4 
 A^^**1*GF  ''r
3  ~~,,SU,K/;<+ #%"$ F^^**6e*LF<<>S(++226":>++226":> ! dnn- ~~,,SU,K/;<+ !NN77 --ce-L1>M"-D ..///N5Dob1$  / /2e!**, 8<;? ##%  /335r8   c           
         ddl }|j                  d      }d}t        | j                        D ]  }	 | j                  j                  |g      }|j                  |      }|r|t        |j                  d            }d|cxk  r	t        k  rn n| j                  j                  |       n4|dz  }| j                  r#t        j                  d| d| dt         d        |dkD  r t        j                  d	| d
t         d       t        | j                        dk(  rt        j                   dt         d       y| j                  r4t        j                  d	t        | j                         dt         d       yy# t        $ r Y Yw xY w)z
        Precompute audio code token IDs (tokens matching <|audio_code_\d+|>).
        These tokens should be blocked during caption generation.
        Only tokens with code values in range [0, MAX_AUDIO_CODE] are included.
        r   N^<\|audio_code_(\d+)\|>$rP   zSkipping audio code token z with invalid code value z (max: )zFound z7 audio code tokens with values outside valid range [0, ]z:No valid audio code tokens found in vocabulary (range [0, z]). Code generation may fail.z$ valid audio code tokens (range [0, z]))recompiler   r   r<   decodematchintgroupMAX_AUDIO_CODEr   addr>   r   	Exceptionr   warning)r   r   audio_code_patterninvalid_tokens_counttoken_id
token_textr   
code_values           r9   r   z@MetadataConstrainedLogitsProcessor._precompute_audio_code_tokens$  s    	ZZ(CD  doo.H!^^22H:>
*00<!$U[[^!4JJ8.81155h?,1,::"LL+EhZOhishtt{  }K  |L  LM  *N  O /"  !#LL6"6!77no}n~~  A  B t(()Q.NNWXfWg  hE  F  GZZLL6#d&?&?"@!AAefteuuwxy   s   B+E::	FFr   c                     ddl }|j                  d      }	 | j                  j                  |g      }|j	                  |      }|rt        |j                  d            S 	 y# t        $ r Y yw xY w)z
        Extract audio code value from a token ID.
        
        Args:
            token_id: Token ID to extract code value from
            
        Returns:
            Code value if token is a valid audio code token, None otherwise
        r   Nr   rP   )r   r   r<   r   r   r   r   r   )r   r   r   r   r  r   s         r9   _extract_code_from_tokenz;MetadataConstrainedLogitsProcessor._extract_code_from_tokenI  sy     	ZZ(CD	..z:J&,,Z8E5;;q>** 
   		s   AA" "	A.-A.c                 0   | j                   sd| _        d| _        yt        j                  d| j
                  t        j                        }t        | j                         }t        d      |d|f<   || _        t        j                  d| j
                  ft        d      t        j                        }d|d|f<   | j                  d|d| j                  f<   || _        | j                  r-t        j                  dt        | j                          d       yy)a  
        Build a precomputed mask tensor for blocking audio code tokens.
        This mask can be added to scores in O(1) time instead of O(n) loop.
        
        The mask is [1, vocab_size] tensor with -inf at audio code token positions.
        
        Also builds the inverse mask (non_audio_code_mask) for CODES_GENERATION state,
        which blocks all non-audio-code tokens.
        NrP   )dtype-infr   zBuilt audio code masks for z tokens)r   r   r   torchzerosr   float32listfloatfullr   r>   r   r   )r   maskaudio_code_indicesinverse_masks       r9   r   z9MetadataConstrainedLogitsProcessor._build_audio_code_mask`  s     ((#'D '+D$ {{1dooU]]C "$";";< ',FmQ""## zz1doo"6fU]][./Q**+ (12LD---.#/ ::LL6s4;T;T7U6VV]^_ r8   scoresallowed_tokensc                 
   |s|j                  t        d             yt        j                  ||j                  t        j
                        }|d|f   j                         }|j                  t        d             ||d|f<   y)a  
        Apply whitelist constraint inplace: only allow specified tokens, block all others.
        
        This is more efficient than creating a mask tensor because:
        1. No memory allocation for mask
        2. No tensor addition operation
        
        Args:
            scores: [1, vocab_size] scores tensor to modify inplace
            allowed_tokens: List of token IDs to allow (all others will be set to -inf)
        r  Ndevicer  r   )fill_r  r  tensorr  longclone)r   r  r  allowed_indicessaved_valuess        r9   _apply_whitelist_inplacez;MetadataConstrainedLogitsProcessor._apply_whitelist_inplace  st     LLv'  ,,~fmmSXS]S]^a01779 	U6]# &2q/!"r8   .c           
         i }d}d}| j                   j                  |d      }| j                  rD|D cg c]  }| j                   j                  |g        }}t	        j                  d| d|        | j
                  D ]  }||z   }| j                   j                  |d      }	d}
t        |	      t        |      k\  r|	dt        |       |k(  rt        |      }
|
&| j                  rt	        j                  d| d	       |	|
d }|s&| j                  rt	        j                  d
| d	       |d   }| j                   j                  |g      }|j                         r!|j                         d   j                         nd}|dvr-| j                  rt	        j                  d| d| d| d       5t        t        |      dz         D ]t  }t        |d|       }||vrt               ||<   |t        |      k  r||   }||   j                  |       J| j                  sW||   j                  | j                         v  | j                  rt	        j                  dt        |       d       t               }||v rZ||   }t        |      D cg c])  }|t!        | j                   j                  |g            f+ }}t	        j                  d|        |S c c}w c c}w )a  
        Build keyscale prefix to allowed tokens mapping based on ACTUAL tokenization.
        
        IMPORTANT: Uses token ID sequences as keys, NOT strings, to avoid tokenization mismatches.
        
        CRITICAL FIX: The tokenizer may merge the context's trailing space into the next token.
        For example:
        - "keyscale: " tokenizes to [10563, 2246, 25, 220] -> ['keys', 'cale', ':', ' ']
        - "keyscale: G major" tokenizes to [10563, 2246, 25, 479, 3598] -> ['keys', 'cale', ':', ' G', ' major']
        The space ' ' (220) is merged into ' G' (479), so we can't use simple slicing.
        
        Strategy:
        1. For each keyscale (e.g., "G major"), encode the FULL string "keyscale: G major"
        2. Tokenize to get: [10563, 2246, 25, 479, 3598] -> ['keys', 'cale', ':', ' G', ' major']
        3. Find where context prefix ends by matching token sequences (handling space merging)
        4. Extract keyscale value tokens: [479, 3598] (for "G major")
        5. Build prefix tree using token ID sequences as keys
        
        This ensures we get the exact tokenization that occurs during generation.
        rY   
keyscale: Fr   z.Context for matching 'keyscale:' tokenizes to  -> N7Could not find context prefix in full tokenization of '', skippingz"No tokens extracted for keyscale 'r   rL   ABCDEFGzSkipping keyscale 'z': first token is 'z' (id=z), not a noterP   z Built keyscale prefix tree with  token sequence prefixesz%First tokens allowed (empty prefix): )r<   r   r>   r   r   r   r   r   lstripupperr   tupler   r   r   sortedrepr)r   prefix_to_tokensrS   rT   context_token_idstcontext_tokens_strrG   	full_textfull_token_idscontext_end_idxkeyscale_token_idsfirst_token_idfirst_token_str
first_charr   token_prefixnext_token_idempty_prefixfirst_tokensdecoded_firsts                        r9   r   z>MetadataConstrainedLogitsProcessor._build_keyscale_prefix_tree  s   * =? '2#*6' !NN112Mbg1h::FW!XFW$.."7"7"<FW!XLLIJ[I\\`as`tuv ,,H7(BI!^^229QV2WN
 #O >"c*;&<<!"93'8#9:>OO&)*;&<O&::NN%\]f\ggr#st "00@!A &::NN%GzQ\#]^ 02N"nn33^4DEO@O@V@V@X//1!4::<^`J*::LL#6xj@STcSddjkyjz  {H  "I  J 312Q67$%7%;<'7758U$\2s-..$6q$9M$\266}E ))(6::4;M;MN 8W -v ::LL;C@P<Q;RRjkl 7L///=PVWcPd ePd1!T$..*?*?*D%E!FPd eD]OTUQ "YJ !fs   #K.K#rO   rS   rT   c                    i }|r| j                   j                  |d      ng }|D ]  }||z   }| j                   j                  |d      }d}	t        |      t        |      k\  r|dt        |       |k(  rt        |      }	|	&| j                  rt	        j
                  d| d       ||	d }
t        t        |
      dz         D ]t  }t        |
d|       }||vrt               ||<   |t        |
      k  r|
|   }||   j                  |       J| j                  sW||   j                  | j                         v  |S )a  
        Build prefix tree for numeric field based on actual tokenization with context.
        
        IMPORTANT: Uses token ID sequences as keys, NOT strings, to avoid tokenization mismatches.
        
        Args:
            valid_values: List of valid numeric strings (e.g., ["30", "31", ..., "300"])
            context_prefix_for_matching: Context string that state machine generates (e.g., "bpm:") - no space
            context_prefix_for_tokenization: Context string for tokenization (e.g., "bpm: ") - with space
            
        Returns:
            Dict mapping token ID sequence prefix -> set of allowed token IDs
        Fr   Nr   r!  rP   )r<   r   r   r>   r   r   r   r&  r   r   r   )r   rO   rS   rT   r)  r*  	value_strr-  	token_idsr/  value_token_idsr   r4  r5  s                 r9   r   z=MetadataConstrainedLogitsProcessor._build_numeric_prefix_tree  sn   & =? mHDNN112Mbg1h  NP &I7)CI--iE-RI #O9~%6!774c"3459JJ&)*;&<O&::NN%\]f\ggr#st ((89O 3/!34$_Ra%89'7758U$\2s?++$3A$6M$\266}E ))(6::4;M;MN 5) &H  r8   c           	         i }d}d}| j                   j                  |d      }| j                  rD|D cg c]  }| j                   j                  |g        }}t	        j                  d| d|        | j
                  D ]=  }||z   }| j                   j                  |d      }	d}
t        |	      t        |      k\  r|	dt        |       |k(  rt        |      }
|
&| j                  rt	        j                  d| d	       |	|
d }|s&| j                  rt	        j                  d
| d	       t        t        |      dz         D ]t  }t        |d|       }||vrt               ||<   |t        |      k  r||   }||   j                  |       J| j                  sW||   j                  | j                         v @ | j                  rt	        j                  dt        |       d       t               }||v rZ||   }t        |      D cg c])  }|t        | j                   j                  |g            f+ }}t	        j                  d|        |S c c}w c c}w )a   
        Build language prefix to allowed tokens mapping based on ACTUAL tokenization.
        Similar to keyscale prefix tree but for language codes.
        
        Uses token ID sequences as keys, NOT strings, to avoid tokenization mismatches.
        rZ   
language: Fr   z.Context for matching 'language:' tokenizes to r  Nr   r!  z"No tokens extracted for language 'rP   z Built language prefix tree with r#  z2First tokens allowed for language (empty prefix): )r<   r   r>   r   r   r   r   r   r   r&  r   r   r   r'  r(  )r   r)  rS   rT   r*  r+  r,  langr-  r.  r/  lang_token_idsr   r4  r5  r6  r7  r8  s                     r9   r   z>MetadataConstrainedLogitsProcessor._build_language_prefix_treeS  sv    =?&1#*6' NN112Mbg1h::FW!XFW$.."7"7"<FW!XLLIJ[I\\`as`tuv((D7$>I!^^229QV2WN"O>"c*;&<<!"93'8#9:>OO&)*;&<O&::NN%\]f\ggr#st+O,<=N!::NN%Gv[#YZ3~.23$^BQ%78'7758U$\2s>**$21$5M$\266}E))(6::4;M;MN 4+ )D ::LL;C@P<Q;RRjkl 7L///=PVWcPd ePd1!T$..*?*?*D%E!FPd eQR_Q`ab[ "YT !fs   #I.Ic                 h   t        d       t        d       t        d       d| j                  v rn| j                  d   }t        dt        |       d       t        |      D ]8  }| j                  j                  |g      }t        d| dt        |              : nt        d       g d	}|D ]  }||d
z   fD ]  }|| j                  v s| j                  |   }t        dt        |       dt        |       d       t        |      D ]8  }| j                  j                  |g      }t        d| dt        |              :   t        dt        | j                                t        t        | j                              dd }|D ]  }	t        dt        |	               t        d       y)z
        Diagnose the keyscale prefix tree to help debug generation bias.
        Call this method to print detailed information about allowed tokens at each prefix.
        z<============================================================zKEYSCALE PREFIX TREE DIAGNOSISrL   z&
[Empty prefix] Allowed first tokens (z total):z  Token z: z#
WARNING: Empty prefix not in tree!)ABCDEFGr   z	
[Prefix z] Allowed tokens (z):z
[Valid keyscales] Total: Nr   z  )	printr   r   r'  r<   r   r(  r   r  )
r   r7  r+  decodedtest_prefixesr   test_keyr   samplekss
             r9   diagnose_keyscale_prefix_treez@MetadataConstrainedLogitsProcessor.diagnose_keyscale_prefix_tree  s   
 	h./h ***44R8L;C<M;NhWXL)..//42d7m_56 * 89 <#F#Vc\2t888!66x@FJtH~&66HVUWXY#F^"&.."7"7"<2d7m_=> ,	 3 $ 	+C0D0D,E+FGHT1123CR8BBtBxj/"  	hr8   c                 <   t         j                  j                  | j                        s/| j                  r"t        j                  d| j                          y	 t         j                  j                  | j                        }|| j                  k  ryt        | j                  dd      5 }g }|D ]F  }|j                         }|s|j                  d      r(|j                  |j                                H || _        || _        | j                          | j                  r8t        j                  dt        | j                         d| j                          ddd       y# 1 sw Y   yxY w# t         $ r"}t        j"                  d	|        Y d}~yd}~ww xY w)
z
        Load genres vocabulary from file. Supports hot reload by checking file mtime.
        File format: one genre per line, lines starting with # are comments.
        zGenres vocab file not found: Nrzutf-8)encodingr   zLoaded z genres from zFailed to load genres vocab: )rq   rr   existsr?   r>   r   getmtimerx   openr   
startswithr   r   rw   _build_genres_trier   r   r   )r   mtimer   rJ   r   es         r9   r   z5MetadataConstrainedLogitsProcessor._load_genres_vocab  sG   
 ww~~d445zz<T=S=S<TUV	@GG$$T%;%;<E///d,,cGDD::<DDOOC$8djjl3 
 %+!*/''')::LL73t/@/@+A*B-PTPfPfOg!hi EDD  	@NN:1#>??	@sI   8E0 E0 +E$E$BE$E0 $E-)E0 -E0 0	F9FFc                     i | _         | j                  D ](  }| j                   }|D ]  }||vri ||<   ||   } d|d<   * | j                  r-t        j                  dt	        | j                         d       yy)z
        Build a trie (prefix tree) from genres vocabulary for efficient prefix matching.
        Each node is a dict with:
          - '_end': True if this node represents a complete genre
          - other keys: next characters in the trie
        T_endzBuilt genres trie with z entriesN)ry   rw   r>   r   r   )r   genrenodechars       r9   rW  z5MetadataConstrainedLogitsProcessor._build_genres_trie  s     &&E##Dt#!#DJDz   DL ' ::LL23t7H7H3I2J(ST r8   c                 .   |r| j                   sy|j                         }t               }ddl}|j	                  d|      }|D cg c]@  }|j                         st        |j                               dk\  s1|j                         B }}|D ])  }| j                  |      }|| j                  |||       + t        | j                         }	|D ]  }||	v s|j                  |        |s"| j                  rt        j                  d       yt        |      | _        i | _        |D ](  }
| j                  }|
D ]  }||vri ||<   ||   } d|d<   * | j                  r2t        j                  dt        |       d	t        |      dd
  d       yyc c}w )a  
        Extract genres from the user's caption that match entries in the vocabulary.
        This creates a smaller trie for faster and more relevant genre generation.
        
        Strategy (optimized - O(words * max_genre_len) instead of O(vocab_size)):
        1. Extract words/phrases from caption
        2. For each word, use trie to find all vocab entries that START with this word
        3. Build a separate trie from matched genres
        Nr   z[,\s\-_/\\|]+   z.No genres matched in caption, using full vocabTr[  zMatched z genres from caption:    z...)rw   r   r   r   r   r   r   _get_genres_trie_node_collect_complete_genresr   r>   r   r  r{   rz   )r   rE   caption_lowermatched_genresr   wordswwordr]  
genres_setr\  r^  s               r9   _extract_caption_genresz:MetadataConstrainedLogitsProcessor._extract_caption_genres  s    d// 	)=9$)OEqQWWY3qwwy>Q;NEO D--d3D--dD.I  **+
Dz!""4(  zzMO '+>&:##% #E++Dt#!#DJDz   DL $ ::LL8C$7#88NtTbOcdfefOgNhhklm E Ps   FF7Fr]  r   result	max_depthc                     |dk  ry|j                  dd      r|j                  |       t        |      dk\  ry|j                         D ]$  \  }}|dvs| j	                  |||z   ||dz
         & y)z}
        Recursively collect all complete genres under a trie node.
        Limited depth to avoid too many matches.
        r   Nr[  Fd   r[  _tokensrP   )getr   r   r   rc  )r   r]  r   rk  rl  r^  
child_nodes          r9   rc  z;MetadataConstrainedLogitsProcessor._collect_complete_genres$  sv    
 >88FE"JJv v;# $

D*..--j&4-QZ]^Q^_ !-r8   c                 4   i | _         i | _        t        | j                        D ](  }	 | j                  j                  |g      }|s$|j                         }|j                         r|j                         }nd}|| j                  |<   |d   j                         }|| j                   vrt               | j                   |<   | j                   |   j                  |       |j                         }|r[||k7  rV|d   j                         }|| j                   vrt               | j                   |<   | j                   |   j                  |       + | j                  r-t        j                  dt        | j                          d       yy# t        $ r Y rw xY w)a  
        Precompute mapping from characters to token IDs and token decoded texts.
        This allows O(1) lookup instead of calling tokenizer.encode()/decode() at runtime.
        
        Time complexity: O(vocab_size) - runs once during initialization
        
        Note: Many subword tokenizers (like Qwen) add space prefixes to tokens.
        We need to handle both the raw first char and the first non-space char.
        r   r   z$Precomputed char->token mapping for z unique charactersN)r|   _token_to_textr   r   r<   r   r   r   rstripr   r   r$  r   r>   r   r   )r   r   text
text_lowernormalized_textr3  stripped_textfirst_nonspace_chars           r9   r}   zAMetadataConstrainedLogitsProcessor._precompute_char_token_mapping7  s}    02.0 doo.H ~~,,hZ8
 "ZZ\
##%&0&7&7&9O&)O0?##H- "!W]]_
T%9%997:uD((4$$Z044X> !% ]d%:*7*:*@*@*B'*$2F2FFDGE,,-@A(()<=AA(K= /F ::LL?DDXDX@Y?ZZlmn   s   F
DF

	FFc                 
   t         j                  j                  | j                        sy	 t         j                  j	                  | j                        }|| j
                  kD  r| j                          yy# t        $ r Y yw xY w)zDCheck if genres vocab file has been updated and reload if necessary.N)rq   rr   rS  r?   rT  rx   r   r   )r   rX  s     r9   _try_reload_genres_vocabz;MetadataConstrainedLogitsProcessor._try_reload_genres_vocabk  sk    ww~~d445	GG$$T%;%;<Et...'') / 		s   AA6 6	BBc                 ^    | j                   }|j                         D ]  }||vr y||   } |S )z
        Get the trie node for a given prefix.
        Returns None if the prefix is not valid (no genres start with this prefix).
        N)ry   r   )r   r   r]  r^  s       r9   rb  z8MetadataConstrainedLogitsProcessor._get_genres_trie_nodew  s:    
 LLND4:D # r8   rv  c                 p    | j                  |j                               }|duxr |j                  dd      S )z>Check if the given text is a complete genre in the vocabulary.Nr[  F)rb  r   rq  )r   rv  r]  s      r9   _is_complete_genrez5MetadataConstrainedLogitsProcessor._is_complete_genre  s3    ))$**,74;DHHVU$;;r8   triec                 J    |}|j                         D ]  }||vr y||   } |S )zGGet a trie node from a specific trie (helper for caption vs full trie).N)r   )r   r  r   r]  r^  s        r9   _get_trie_node_from_triez;MetadataConstrainedLogitsProcessor._get_trie_node_from_trie  s2    LLND4:D # r8   c                    | j                   sg S | j                  j                         }|j                         }d}d}| j                  r4|dk(  r| j                  }d}n | j                  | j                  |      }|d}|#|dk(  r| j                  }n| j                  |      }|| j                  r| j                  gS g S t        d |j                         D              }|j                  dd      }|s>t               }|r'| j                  r|j                  | j                         t        |      S t               }|D ]/  }	|	| j                  v s|j                  | j                  |	          1 |r| j                  n| j                  }
t               }|D ]  }| j                   j                  |d      }|r|j                         sd|v sd|v r|j                  |       K|j#                  d      s|j#                  d      r||z   }n||z   }| j                  |
|      }||j                  |        |r'| j                  r|j                  | j                         t        |      S )	a  
        Get allowed tokens for genres field based on trie matching.
        
        The entire genres string (including commas) must match a complete entry in the vocab.
        For example, if vocab contains "pop, rock, jazz", the generated string must exactly
        match that entry - we don't treat commas as separators for individual genres.
        
        Strategy:
        1. If caption-matched genres exist, use that smaller trie first (faster + more relevant)
        2. If no caption matches or prefix not in caption trie, fallback to full vocab trie
        3. Get valid next characters from current trie node
        4. For each candidate token, verify the full decoded text forms a valid trie prefix
        FNrL   Tc              3   *   K   | ]  }|d vs|  yw)ro  Nr7   ).0r   s     r9   	<genexpr>zPMetadataConstrainedLogitsProcessor._get_allowed_genres_tokens.<locals>.<genexpr>  s     ^*=QJ]A]q*=s   	r[  r   r   )rw   rh   r   r   rz   r  ry   rb  r   r   keysrq  r   r  r|   updatert  rV  )r   accumulatedcurrent_genre_prefixuse_caption_triecurrent_nodevalid_next_charsis_completeallowedcandidate_tokensr^  active_trier   decoded_normalized
new_prefixnew_nodes                  r9   _get_allowed_genres_tokensz=MetadataConstrainedLogitsProcessor._get_allowed_genres_tokens  so      I ,,224*002 ! ###r)#77#' #<<T=U=UWkl+'+$ #r)#//#99:NO!!**++I ^,*;*;*=^^ #&&vu5eGt11D../=  5$Dt+++ ''(<(<T(BC %
 3Cd..HXHX %(H!%!4!4!8!82!F%-?-E-E-G**c5E.EKK) ",,S15G5R5RSV5W14FF
14FF
 44[*MH#H%+ )0 4--KK**+G}r8   c                     t         j                  | _        d| _        d| _        g | _        d| _        g | _        d| _        d| _	        d| _
        d| _        d| _        y)z/Reset the processor state for a new generation.r   rL   NF)r   r   rf   rg   rh   ri   rc   rn   ro   rj   rk   rl   rm   r   s    r9   resetz(MetadataConstrainedLogitsProcessor.reset  s_    ''
!"!#%'"&(#"&%*"#$ #"$r8   rF   c                     || _         |L|dkD  rGt        |dz        | _        | j                  r't	        j                  d| d| j                   d       yyd| _        | j                  rt	        j                  d       yy)z
        Set the target duration for codes generation.
        
        Args:
            duration: Target duration in seconds. If None, no duration constraint is applied.
                     5 codes = 1 second, so target_codes = duration * 5.
        Nr   ra  zSet target duration: s -> z codesz/Target duration cleared, no duration constraint)ra   r   rb   r>   r   )r   rF   s     r9   set_target_durationz6MetadataConstrainedLogitsProcessor.set_target_duration  s|      (HqL #HqL 1Dzz4XJeDDUDUCVV\]^  !%DzzNO r8   c           
         || j                   k(  ry| j                   }|| _         || j                  d   d<   t        | j                  d   d   | j                  d   d   dz         D cg c]  }t        |       c}| _        | j                  | j                  dd      | _        | j                  r3t        j                  d	| d
| dt        | j                         d       yyc c}w )a  
        Dynamically update the maximum allowed duration for constrained decoding.
        
        This method should be called when GPU configuration changes (e.g., LM initialization state changes).
        It rebuilds the duration prefix tree to constrain duration values to the new maximum.
        
        Args:
            max_duration: Maximum duration in seconds (e.g., 120 for 2 minutes, 360 for 6 minutes)
        NrF   rN   rM   rP   rU   rV   rR   zUpdated max duration: r  zs, rebuilt prefix tree with z values)
rA   r~   r   r   r   r   r   r>   r   r   )r   rA   old_maxr   s       r9   set_max_durationz3MetadataConstrainedLogitsProcessor.set_max_duration  s8    4,,,##( /;$U+ 7<D<L<LZ<XY^<_aeaqaqr|a}  D  bE  HI  bI  7J  &K  7Jc!f  7J  &K" %)$C$C&&(3,8 %D %
! ::LL1'%~Mijmnr  oI  oI  kJ  jK  KR  S  T  &Ks   'C$	fixed_strc                    || j                   d }|sg S | j                  r:t        j                  dt        |       d| j                    dt        |              d}d}t	        t        |      dd      D ]  }|d| }| j                  j                  |d      }|s(t        |      d	k(  s7|d   }|}| j                  rJt        j                  d
t        |       d| dt        | j                  j                  |g                     n ||gS i }t	        d	t        t        |      d	z   d            D ]  }|d| }| j                  j                  |d      }|s(|d   }	| j                  j                  |	g      }
|j                         j                         }|
j                         j                         }|j                  |      s|j                  |      s|	|vs	|||	   kD  s|||	<    t        |j                         d d      }|r|D cg c]  \  }}|	 c}}ng }| j                  rt        j                  dt        |       d|dd D cg c])  }|t        | j                  j                  |g            f+ c}        |r:t        j                  dt        |       d| j                    dt        |              |S c c}}w c c}w )a:  
        Get the token IDs that can continue the fixed string from current position.
        Returns list of allowed token IDs.
        
        Strategy: Find the longest prefix that encodes to a single token, and return that token.
        This ensures we generate by tokens, not character-by-character.
        Nz0_get_allowed_tokens_for_fixed_string: fixed_str=z, position_in_state=z, remaining=r   r   Fr   rP   z!Found single-token match: prefix=z, token_id=z, token_text=   c                     | d   S )NrP   r7   )xs    r9   <lambda>zYMetadataConstrainedLogitsProcessor._get_allowed_tokens_for_fixed_string.<locals>.<lambda>u  s    QqTr8   T)keyreversezFallback: returning z	 tokens: ra  zFixed string: z, position: z, remaining: )rg   r>   r   r(  r   r   r<   r   r   rM   r$  r   rV  r'  r   )r   r  	remaining
best_tokenbest_prefix_lenendr   r   r  first_tokendecoded_tokennormalized_prefixnormalized_decodedsorted_tokenstoken_rk  r+  s                     r9   $_get_allowed_tokens_for_fixed_stringzGMetadataConstrainedLogitsProcessor._get_allowed_tokens_for_fixed_string7  sZ    d4456	I::LLKDQZOK\\pqu  rH  rH  qI  IU  VZ  [d  Ve  Uf  g  h 
 YB/Ct_F^^**6e*LF#f+* $AY
"%::LL#DT&\NR]^h]iivw{  }A  }K  }K  }R  }R  T^  S_  }`  xa  wb  "c  d 0 !< CI 2B78Ct_F^^**6e*LF$Qi $ 5 5{m D$*MMO$9$9$;!%2%9%9%;%A%A%C" &001BCGXGcGcdvGw".8C.Q\B]<]69{3 9& ~335>SWX:G6HE1%6R::LL/F}Iv|}~  wA  GB  wAqr4PTP^P^PePeghfiPjKkGl  wA  GB  FC  D  E~d9o->l4KaKaJbboptu~p  pA  B  C 7 GBs   K
.Kmin_valmax_valc                 d   | j                   spt               }t        ||dz         D ](  }|j                  t	        t        |      d                * |D cg c]   }|| j                  v s| j                  |   " c}S t	        | j                         }g }t        d      D ]Z  }t	        | j                   t        |      z         }||kD  r*||k\  r|j                  |       A|dz  |k  sJ|j                  |       \ |D cg c]   }|| j                  v s| j                  |   " c}S c c}w c c}w )z
        Get allowed digit tokens based on accumulated value and range constraints.
        Uses early-blocking to prevent out-of-range values.
        rP   r   r   )rh   r   r   r   r   r   r   r   )	r   r  r  allowed_digitsr   r   currentr  	new_values	            r9   _get_allowed_digit_tokensz<MetadataConstrainedLogitsProcessor._get_allowed_digit_tokens  s"   
 %% UN7GaK0""3s1vay>2 12@[.QAIZIZDZD%%a(.[[d,,-rAD22SV;<I 7"
 G#q!R7*q!% ( /6Pgd>O>O9O!!!$gPP3 \2 Qs   D()D( D-D-prefix_treec                 T    t        | j                        }||v rt        ||         S g S )a  
        Get allowed tokens for numeric field using the precomputed prefix tree.
        
        IMPORTANT: Uses token ID sequence as key (not string) to avoid tokenization mismatches.
        
        Args:
            prefix_tree: Precomputed prefix tree mapping token ID sequence -> set of allowed token IDs
            
        Returns:
            List of allowed token IDs for current accumulated_token_ids
        )r&  ri   r  )r   r  r4  s      r9   _get_allowed_numeric_tokensz>MetadataConstrainedLogitsProcessor._get_allowed_numeric_tokens  s3     T778;&L122 	r8   logitsc                    | j                   syt        | j                         }||k  s||kD  ryt        j                  |d      | j                  rd| j                  f   j                         nd}| j                  ||      }|syt        fd|D              }| j                  rt        j                  d|dd	|d       ||kD  S )
z
        Determine if we should end the current numeric field.
        Returns True if P(newline) > P(any valid digit) AND current value is valid.
        Fr   dimr   Tc              3   H   K   | ]  }d |f   j                           yw)r   N)item)r  r+  probss     r9   r  zOMetadataConstrainedLogitsProcessor._should_end_numeric_field.<locals>.<genexpr>  s#     HAU1a4[--/s   "z%Numeric field decision: newline_prob=z.4fz, max_digit_prob=)
rh   r   r  softmaxr   r  r  rN   r>   r   )	r   r  r  r  r  newline_probr  max_digit_probr  s	           @r9   _should_end_numeric_fieldz<MetadataConstrainedLogitsProcessor._should_end_numeric_field  s    
 %%d,,-W' 1 f"->B>P>PuQ 2 22388:VW 77IHHH::LL@c@RRcdrsvcwxyn,,r8   c                 f   | j                   j                         syt        j                  |d      }| j                  r|d| j                  f   j                         nd}|j                         }| j                  rd|d| j                  f<   |d   j                         j                         }||kD  S )z
        Determine if we should end a text field (genres).
        Returns True if P(newline) > P(any other token) AND we have some content.
        Fr   r  r   )rh   r   r  r  r   r  r  rN   )r   r  r  r  masked_probsmax_other_probs         r9   _should_end_text_fieldz9MetadataConstrainedLogitsProcessor._should_end_text_field  s    
 %%++-f"->B>P>PuQ 2 22388:VW {{}23LD.../%a,,.335n,,r8   c                 |    t        | j                        }|| j                  v rt        | j                  |         S g S )z
        Get allowed tokens for keyscale field using the precomputed prefix tree.
        Uses token ID sequence as key (not string) to avoid tokenization mismatches.
        )r&  ri   r   r  r   r4  s     r9   _get_allowed_keyscale_tokensz?MetadataConstrainedLogitsProcessor._get_allowed_keyscale_tokens  s?     T778444411,?@@ 	r8   c                     t        | j                        }|| j                  v r| j                  | j                  |   v S y)z
        Check if keyscale value is complete and valid.
        Uses token ID sequence to check if current prefix allows newline.
        F)r&  ri   r   r   r  s     r9   _is_keyscale_completez8MetadataConstrainedLogitsProcessor._is_keyscale_complete  sA    
 T7784444%%)B)B<)PPPr8   c                 |    t        | j                        }|| j                  v rt        | j                  |         S g S )z
        Get allowed tokens for language field using the precomputed prefix tree.
        Uses token ID sequence as key (not string) to avoid tokenization mismatches.
        Similar to keyscale.
        )r&  ri   r   r  r  s     r9   _get_allowed_language_tokensz?MetadataConstrainedLogitsProcessor._get_allowed_language_tokens  s?     T778444411,?@@ 	r8   c                 |    t        | j                        }|| j                  v rt        | j                  |         S g S )z
        Get allowed tokens for timesignature field using the precomputed prefix tree.
        Uses token ID sequence as key (not string) to avoid tokenization mismatches.
        )r&  ri   r   r  r  s     r9   _get_allowed_timesig_tokensz>MetadataConstrainedLogitsProcessor._get_allowed_timesig_tokens  s?    
 T778433300>?? 	r8   	input_idsc                 4   | j                   s| j                  |      S | j                  t        j                  k(  r| j
                  dk(  r| j                  | j                  j                  |j                  k7  s#| j                  j                  |j                  k7  r6| j                  j                  |j                  |j                        | _        || j                  z   }| j                  |      S | j
                  dk(  rk| j                  t        j                  k(  rN| j                  |      r=t        j                  | _        d| _        | j                  rt        j                  d       | j                  t        j                  k(  r| j                   | j                   j                  |j                  k7  s#| j                   j                  |j                  k7  r6| j                   j                  |j                  |j                        | _        || j                   z   }| j"                  | j$                  | j                  | j"                  k  rYt'        d      |dd| j$                  f<   | j                  rt        j                  d| j                   d	| j"                   d
       n|dd| j$                  f   j)                         }|j+                  t'        d             ||dd| j$                  f<   | j                  r0t        j                  d| j                   d	| j"                   d       | j                  |      S |j,                  d   }t/        |      D ]%  }| j1                  ||   |||dz          }|d   ||<   ' | j                  |      S )aU  
        Apply constrained decoding by modifying logits.
        
        Args:
            input_ids: [batch_size, seq_len] input token IDs
            scores: [batch_size, vocab_size] logits for next token
            
        Returns:
            Modified scores with invalid tokens masked to -inf and temperature scaling applied
        r   Nr  r   r   zECodes phase: detected </think> in input, skipping to CODES_GENERATIONr  zCodes generation: /z, blocking EOSz, forcing EOSrP   )r=   _apply_temperature_scalingrf   r   r6   re   r   r  r  tor   _input_contains_think_end_tagr5   rc   r>   r   r   rb   r   r  r  r  shaper   _process_single_sequence)r   r  r  
eos_scores
batch_sizer   rk  s          r9   __call__z+MetadataConstrainedLogitsProcessor.__call__   s    ||226::::+++$$49M9M9Y''..&--?4CWCWC]C]agamamCm+/+?+?+B+B&--_e_k_k+B+lD($"6"66226::   G+

h>P>P0P11)<%66
#$ ::LL!hi::222 ''3++22fmmCtG_G_GeGeioiuiuGu/3/G/G/J/JRXR_R_gmgsgs/J/tD,$":"::   ,1B1B1N##d&7&7738=F1d///0zz'9$:J:J9K1TM^M^L__m%no "(4+<+<(<!=!C!C!EJLLv/3=F1d///0zz'9$:J:J9K1TM^M^L__l%mn226::\\!_
 z"A229Q<!A#OFq	F1I #
 ..v66r8   c                 "   | j                   j                  dd      }|syt        |j                  d         D ]T  }||   j	                         }t        t        |      t        |      z
  dz         D ]  }|||t        |      z    |k(  s  y V y)a   
        Check if input contains the </think> closing tag.
        
        Args:
            input_ids: [batch_size, seq_len] input token IDs
            
        Returns:
            True if </think> is found in the input (any sequence in batch)
        r[   Fr   r   rP   T)r<   r   r   r  tolistr   )r   r  think_end_tokensr   seqr   s         r9   r  z@MetadataConstrainedLogitsProcessor._input_contains_think_end_tagi  s      >>00PU0V yq)*AA,%%'C3s8c*:&;;a?@q3/0015EE A + r8   c                     | j                   t        j                  k(  s| j                   t        j                  k(  r| j                  }n| j
                  }||S |dk  rd}||z  S )a  
        Apply temperature scaling based on current generation phase.
        
        Temperature scaling: logits = logits / temperature
        - Lower temperature (< 1.0) makes distribution sharper (more deterministic)
        - Higher temperature (> 1.0) makes distribution flatter (more diverse)
        
        Args:
            scores: [batch_size, vocab_size] logits
            
        Returns:
            Temperature-scaled logits
        r   gư>)rf   r   r5   r6   r`   r_   )r   r  temperatures      r9   r  z=MetadataConstrainedLogitsProcessor._apply_temperature_scaling  sf     ::222djjHDVDV6V00K33K M !K ##r8   
field_namec                    | j                   j                  |      }|ydddddddd	}||   }| | d
}| j                  j                  |d      }|dz   }| j                  j                  |d      }t	        |      t	        |      k\  r|dt	        |       |k(  r|t	        |      d S | j
                  rt        j                  d| d       |S )a  
        Get token sequence for a user-provided field (field_name + value + newline).
        Uses the same tokenization logic as prefix tree building.
        
        Args:
            field_name: Field name ("bpm", "caption", "duration", "keyscale", "language", "timesignature")
            
        Returns:
            List of token IDs for the complete field, or None if field is not provided
        NrQ   z	caption: rV   r  r>  rW   zgenres: rC   rX   Fr   :z(Could not match prefix tokens for field z, using all tokens)r^   rq  r<   r   r   r>   r   r   )	r   r  valuefield_to_prefixr   r-  r   prefix_for_matchingprefix_tokenss	            r9   _get_user_provided_field_tokenszBMetadataConstrainedLogitsProcessor._get_user_provided_field_tokens  s     ++//
;= "$$$. 
 !,hugR(	 &&yU&K )3.--.AV[-\ v;#m,,8K]9K1LP]1]#m,-.. zz!I*UghiMr8   c           
         | j                   r$| j                   d   }| j                  ||g       |S | j                  | j                  v r| j                  | j                     }| j	                  |      }|r| j                  t
        j                  k(  ry| j                  rmt        |      | j                  z
  }|dk  rP| j                  D| j                  || j                  g       | j                  rt        j                  d| d       |S | j                  ||       |S | j                  t
        j                  k(  rX| j                  rL| j                  @| j                  || j                  g       | j                  rt        j                  d       |S | j                  }| j                          | j                  | j                  v rH| j                  r:t        j                  d|j                   d| j                  j                   d       |S |j!                          | j#                  ||      S | j                  t
        j$                  k(  r| j&                  d	   r| j                   sf| j(                  sZ| j&                  d	   }d
| d}	| j*                  j-                  |	d      }
|
r&|
| _         d	| _        | j                  ||
d   g       |S | j1                  | j2                        }t5        | j(                        }|| j2                  v r+| j6                  | j2                  |   v r|| j6                  gz   }| j                  ||       |S | j                  t
        j8                  k(  r	| j&                  d   r| j                   sf| j:                  sZ| j&                  d   }d
| d}	| j*                  j-                  |	d      }
|
r&|
| _         d| _        | j                  ||
d   g       |S | j<                  rut?        j@                  |d         jC                         }| j*                  jE                  |g      }t        |      dkD  r|d   dvrd| _        d| _#        d| _$        |S d| _        | jF                  r|S | jJ                  tM        d      |d| jJ                  f<   | jN                  | jN                  jP                  |jP                  k7  s#| jN                  jR                  |jR                  k7  r6| jN                  jU                  |jP                  |jR                        | _'        || jN                  z   }| jV                  dk\  r+| j6                  | j                  || j6                  g       |S |S | j                  t
        jX                  k(  r| j&                  d   r| j                   sf| j(                  sZ| j&                  d   }d
| d}	| j*                  j-                  |	d      }
|
r&|
| _         d| _        | j                  ||
d   g       |S | jZ                  t]        t_        | jZ                              }t        | j:                        }|t        |      k  r>t_        ||         }|| j`                  v r| j                  || j`                  |   g       |S | j6                  r| j                  || j6                  g       |S | j1                  | jb                        }t5        | j(                        }|| jb                  v r+| j6                  | jb                  |   v r|| j6                  gz   }| j                  ||       |S | j                  t
        jd                  k(  r| j&                  d   r| j                   sf| j:                  sZ| j&                  d   }d
| d}	| j*                  j-                  |	d      }
|
r&|
| _         d| _        | j                  ||
d   g       |S | jg                          | ji                         }|r| j                  ||       |S | jj                  rZ| j6                  r| j                  r#t        j                  d| j:                   d       | j                  || j6                  g       |S | jm                  |      r;| j6                  ro| j                  || j6                  g       | j                          |S | j:                  jo                         s&| j6                  rtM        d      |d| j6                  f<   |S | j                  t
        jp                  k(  r;| j&                  d   r| j                   sf| j(                  sZ| j&                  d   }d
| d}	| j*                  j-                  |	d      }
|
r&|
| _         d| _        | j                  ||
d   g       |S t5        | j(                        }|| jr                  v rF| j6                  | jr                  |   v r+| j6                  rn| j                  || j6                  g       |S | ju                         }|r| j                  ||       |S | j6                  r| j                  || j6                  g       |S | j                  t
        jv                  k(  r| j&                  d   r| j                   sf| j(                  sZ| j&                  d   }d
| d}	| j*                  j-                  |	d      }
|
r&|
| _         d| _        | j                  ||
d   g       |S | j(                  sXt5               }|| jx                  v rt{        | jx                  |         }|rt?        j|                  ||jP                  t>        j~                        }|d|f   }t?        j@                  |      jC                         }||   }| j                  ||g       | j                  r^| j*                  jE                  |g      }t        j                  d| dt        |       dt        |       d       |S | j6                  r| j                  || j6                  g       |S | j6                  r| j                  || j6                  g       |S t5        | j(                        }|| jx                  v rF| j6                  | jx                  |   v r+| j6                  rn| j                  || j6                  g       |S | j                         }|r| j                  ||       |S | j6                  r| j                  || j6                  g       |S | j                  t
        j                  k(  r| j&                  d   r| j                   sf| j(                  sZ| j&                  d   }d
| d}	| j*                  j-                  |	d      }
|
r&|
| _         d| _        | j                  ||
d   g       |S t5        | j(                        }|| j                  v rF| j6                  | j                  |   v r+| j6                  r| j                  || j6                  g       |S | j                         }| j                  ||       |S ) zMProcess a single sequence and return modified scores (inplace when possible).r   r   zIstop_at_reasoning=True: forcing EOS near end of </think> tag (remaining: z chars)zAstop_at_reasoning=True: forcing EOS after completing </think> tagzState transition from z to z+ still in fixed_strings, avoiding recursionrD   r   rX   Fr   rE   z 	TrL   r  r  i   rF   rJ   z!No valid genre continuation for 'z', forcing newlinerG   rH   z%Language field: selected top-1 token z (z) from z candidatesrI   )Ern   r  rf   r   r  r   r4   rd   r   rg   r   r>   r   _transition_to_next_stater   namezero_r  r"   r^   ri   r<   r   ro   r  r   r&  r   r%   rh   rj   r  argmaxr  r   rl   rm   r   r  r   r  r  r  rk   r'   ra   r   r   r   r   r*   r|  r  rw   r  r   r-   r   r  r0   r   r  r  r  r(  r  r2   r   r  )r   r  r  
next_tokenr  r  remaining_chars	old_stater  
value_textvalue_tokensr4  top_token_idtop_token_text
target_strcurrent_pos
next_digitr6  r  candidate_indicescandidate_scoresbest_idxs                         r9   r  z;MetadataConstrainedLogitsProcessor._process_single_sequence  s/    &&44Q7J))&:,?M::+++**4::6I??	JG ::!7!77D<R<R&))nt7M7M&MO&",,,8 99&4CTCTBUV#zz &/x  zI  yJ  JQ  .R  !S#)M --fg>R
 I
 ::!7!77D<R<R((455ft?P?P>QR::"LL+ln% JJ	..0::!3!33zz)?	?OtTXT^T^TcTcSd  eP  (Q  R!M44YGGZZ8---**51=dFaFajn  kE  kE33E: r]
#~~44ZTY4Z2>D/.3D+11&<?:KL!M 66t7K7KLG !!;!;<Lt3338J8JdNbNbcoNp8p!T%7%7$88))&':r o ZZ8111 **95A$JeJenr  oE  oE33I> r]
#~~44ZTY4Z2>D/.7D+11&<?:KL!M ))$||F1I6;;=!%!6!6~!F ~&*~a/@/M 27D.*.D'.0D+!M 27D. ""  "".16vq$---. ##/''..&--?4CWCWC]C]agamamCm+/+?+?+B+B&--_e_k_k+B+lD($"6"66 ''3.%%111&4;M;M:NO!M MZZ8222**:6B4KfKfos  pJ  pJ33J? r]
#~~44ZTY4Z2>D/.8D+11&<?:KL!M ##/ T%9%9!:;
!$"8"89Z0!$Z%<!=J!T%6%6655ft?P?PQ[?\>]^r m ))55ft?Q?Q>RSj c ::4;T;TU  %T%?%?@4#<#<<ASASW[WpWpq}W~A~%););(<<G--fg>T Q ZZ8000**84@IdIdmq  nD  nD33H= r]
#~~44ZTY4Z2>D/.6D+11&<?:KL!M ))+ 557G--fg>d c "" %%zz'HI_I_H``r%st11&4;M;M:NOV Q ..v6))55ft?Q?Q>RS668J E  11779--<A&MF1d&8&8#89@ { ZZ8222**:6B4KfKfos  pJ  pJ33J? r]
#~~44ZTY4Z2>D/.8D+11&<?:KL!M !!;!;<Lt888T=O=OSWSlSlmySz=z%%11&4;M;M:NOT O ;;=11&'BJ C ))55ft?Q?Q>RS@ } ZZ8222 **:6B4KfKfos  pJ  pJ33J? r]
#~~44ZTY4Z2>D/.8D+11&<?:KL!M --$w4#<#<<'+D,E,El,S'T$' -2LL9IRXR_R_glgqgq,r)+1!5F2F+G( $)<<0@#A#F#F#H'7'A 55f|nM::-1^^-B-BL>-RN"LL+PQ]P^^`aeftau`vv}  B  CS  T  ~U  U`  *a  bh c  -- 99&4CUCUBVW` [ ))55ft?Q?Q>RSX Q  %T%?%?@4#<#<<ASASW[WpWpq}W~A~))55ft?Q?Q>RSH C #??AG55fgF> 9  -- 99&4CUCUBVW6 3 ZZ8111**?;GPTPkPktx  uO  uO33OD r]
#~~44ZTY4Z2>D/.=D+11&<?:KL!M !!;!;<Lt777D<N<NRVRjRjkwRx<x%%11&4;M;M:NO  ::<--fg>r8   c                 "   | j                   | j                  v r| j                   }| j                  | j                      }|t        j                  k(  r@| j                  dk(  r1t        j
                  }| j                  rt        j                  d       || _         d| _        d| _	        g | _
        d| _        d| _        d| _        d| _        | j                  r:t        j                  d|j                   d| j                   j                          yyy)	z!Transition to the next FSM state.r   zGgeneration_phase='understand': allowing free-form lyrics after </think>r   rL   FzFSM transition: r  N)rf   r   r   r4   re   r6   r>   r   rg   rh   ri   rj   rk   rl   rm   r  )r   r  r   s      r9   r  z<MetadataConstrainedLogitsProcessor._transition_to_next_state=  s    ::(

I4J H222((L8 "*!3!3Jzz'np $DJ%&D"%'D")+D&).D&'(D$"'D&(D#zz/	/?tDJJOOCTUV 3 )r8   generated_token_idc                    | j                   sy| j                  t        j                  k(  ry| j                  t        j                  k(  r]| xj
                  dz  c_        | j                  r;| j                  /t        j                  d| j
                   d| j                          y| j                  rB| j                  d   }||k7  r4| j                  r(t        j                  d| d| d| j                          | j                  j                  d       | j                  s| j                  r"t        j                  d	| j                          | j                  }d| _        | j                  |      }|rn| j                  }|| _        d| _        d
| _        g | _        | j                  r9t        j                  d|j"                   d| j                  j"                          y| j%                          y| j&                  j)                  |g      }| j                  r;t        j                  dt+        |       d| d| j                  j"                          | j                  | j,                  v r| j,                  | j                     }| xj                  t/        |      z  c_        | j                  t/        |      k\  r| j                  t        j0                  k(  r| j2                  r|| j                  }t        j                  | _        d| _        d
| _        g | _        | j                  rKt        j                  d|j"                   d| j                  j"                          y| j%                          yyy| j                  t        j4                  t        j6                  t        j8                  fv r|| j:                  k(  r5| j                  }| j%                          | j                  | j,                  v r_y| j                   j=                  |       |j?                         jA                         r$| xj                  |j?                         z  c_        yyy| j                  t        jB                  k(  rO|| j:                  k(  r)| j%                          | j                  | j,                  v ry| xj                  |z  c_        yy| j                  t        jD                  k(  r| xjF                  dz  c_#        | xj                  |z  c_        d|v rd| _$        nd| _$        | jJ                  r| xjL                  |z  c_&        d|v s|j?                         dk(  r| jL                  j?                         }|jO                  d      j?                         jQ                         }| j                  r!t        j                  dt+        |              t        j6                  t        jB                  t        jR                  t        jT                  t        j8                  d}	||	v r| j                  }|	|   | _        d| _        d
| _        g | _        d| _%        d
| _&        | j                  rt        j                  d|j"                   d| j                  j"                          y| j                  r"t        j                  dt+        |       d       d| _%        d
| _&        | j%                          yyyy| j                  t        jR                  k(  rj|| j:                  k(  r)| j%                          | j                  | j,                  v r2y| j                   j=                  |       | xj                  |z  c_        yy| j                  t        jT                  k(  rj|| j:                  k(  r*| j%                          | j                  | j,                  v ryy| j                   j=                  |       | xj                  |z  c_        yy)z
        Update internal state after a token has been generated.
        This should be called after each token generation.
        
        Args:
            generated_token_id: The token ID that was just generated
        NrP   zCodes count: r  r   zExpected token z	 but got z for user-provided field z,Completed injection of user-provided field: rL   z-FSM transition (after user field injection): r  zGenerated token: z (id=z	), state=z$FSM transition (stop_at_reasoning): rX   TFr  z#Detected field name after caption: )rF   rJ   rG   rH   rI   z!FSM transition (caption ending): z"Unknown field name after caption: z, forcing transition)+r=   rf   r   r6   r5   rc   r>   rb   r   rn   r   ro   popr   rg   rh   ri   r  r  r<   r   r(  r   r   r4   rd   r"   r'   r2   r   r   r   isdigitr*   r%   rk   rj   rl   rm   ru  r   r-   r0   )
r   r   expected_tokenr  r   r  	token_strr  field_name_fullfield_name_to_value_states
             r9   update_statez/MetadataConstrainedLogitsProcessor.update_state[  s    ||::+++::222!zzd//;}T-=-=,>a@Q@Q?RST &&!88;N!^3::NN_^4DIN`Maaz{  |S  |S  {T  $U  V ''++A. ..::LL#OPTPgPgOh!ij!44
*.' "77
C
 $

I!+DJ-.D*-/D*13D.zz'TU^UcUcTddhimisisixixhy%z{  224NN))+=*>?	::LL,T)_,=UCUBVV_`d`j`j`o`o_pqr::+++**4::6I""c)n4" %%Y7::!7!77D<R<R !%

I!)!3!3DJ-.D*-/D*13D.zz'KINNK[[_`d`j`j`o`o_p%qr224 " 8 ZZH..0G0GI_I_``!T%7%77 JJ	..0
 ::!3!33 **112DE??$,,.**ioo.??* / 4 ZZ8000!T%7%77..0
 ::!3!33 &&)3&	 4 ZZ8111$$)$ ""i/" y -1* .3* ""''94' )#y'8C'?&*&=&=&C&C&EO!0!7!7!<!B!B!D!J!J!LJzz'J4PZK[J\%]^ %-$;$;"*"7"7$,$;$;$,$;$;)1)?)?1- "%>>$(JJ	%>z%J
12.13.572.3+24/::"LL+LY^^L\\`aeakakapap`q)rs  ::"NN-OPTU_P`Oaau+vw.3+24/668 &9 (@	 #T ZZ8222!T%7%77..0
 ::!3!33 **112DE&&)3& 4 ZZ8222!T%7%77..0::!3!33 4 **112DE&&)3& 3r8   )TFNTN)N)rL   rL   )2   )Er   r   r   r   r   boolr   r   r   r   r   r   r   r   r   r   staticmethodr   r   r   r   r   rp   r   r  r   r  Tensorr	   r  r   r   r   r   r   rO  r   rW  rj  r   rc  r}   r|  rb  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  
LongTensorFloatTensorr  r  r  r  r  r  r  r7   r8   r9   r;   r;   Q   s   . +/ &*^( ^( ^( 	^(
 $C=^( ^( sm^(@-&3 -&8H;M -&^)IV(D (
(T (
(d (
 'S 'S ' '@&$ &&# &"X(4Xc]8J3K*L "XHR6n#zJ # .'`R2u|| 2TRUY 2[_ 26n T%S/3s82K-L n f ,./1	= 3i=  &)=  *-	= 
 
eCHos3x'	(= ~< T%S/3s82K-L < |%P@>U(6ns 6np`T `3 ` `X[ `&2oh

C 
HTN 
<s <t <
T 3 8D> dDI dL%PHUO P$TS T@Fc Fd3i FP#Q #Qs #QtCy #QJtE#s(OSQTX<U7V [_`c[d *- -s -UX -]a -<-U\\ -d -&d3i 	t 	d3i T#Y G7##G7 !!G7 
			G7Ru7G7G D 2$1B1B $uGXGX $>,# ,(4PS9BU ,\l##l !!l 
			l\W<G4s G4r8   r;   )"r   enumr   r   typingr   r   r   r   r	   r
   r   logurur   transformersr   &transformers.generation.logits_processr   rq   r  acestep.constantsr   r   r   r   r   r   r   r   r   r   r   r   r;   r7   r8   r9   <module>r     sW   *  B B B  & B 	   " t 8Q#4 Q#4r8   