
    Njl                     b   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Zd dlZd dlm	Z	 d dl
mZ d dlZd dlZd dlZd dlZd dlmZ d dlmZ d dlmZ d dlmZ d dl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% d dl&m'Z' d dl(m)Z) d dl*m+Z+ d dl,m-Z-m.Z. d dl/m0Z0m1Z1m2Z2 d dl3m4Z4m5Z5 d dl6m7Z7m8Z8m9Z9 d dl:m;Z;m<Z<m=Z=m>Z>m?Z?m@Z@ d dlAmBZBmCZC  ej                  d        ej                  eF      ZGd eH e;j                                dZJ G d dej                        ZLdeMdeNdeNdeOeMePf   fdZQdeHeP   d eHeP   deHeP   fd!ZRdeHeP   d"eHeP   deHeOePeMf      fd#ZSd$eMd%ePdeNdeNdeHeM   f
d&ZTd'eUeMeUeMej                  f   f   d(eMez  fd)ZWd*eMez  d+ej                  deUeMeUeMej                  f   f   fd,ZYy)-    N)	lru_cache)Path)nn)
functional)Self)TokenizedText)
audio_read)convert_audio)DEFAULT_EOS_THRESHOLDDEFAULT_LANGUAGEDEFAULT_LSD_DECODE_STEPSDEFAULT_NOISE_CLAMPDEFAULT_TEMPERATUREMAX_TOKEN_PER_CHUNK)FlowLMModel)	MimiModel)mimi_transformer)DummyQuantizer)SEANetDecoderSEANetEncoder)StatefulModuleincrement_stepsinit_states)RECOMMENDED_CONFIGapply_dynamic_int8)CONFIGS_DIRConfigload_config)_ORIGINS_OF_PREDEFINED_VOICES
DEBUG_MIMIdisplay_execution_timedownload_if_necessaryget_predefined_voicesize_of_dict)get_flow_lm_state_dictget_mimi_state_dict   zWe could not download the weights for the model with voice cloning, but you're trying to use voice cloning. Without voice cloning, you can use our catalog of voices z. If you want access to the model with voice cloning, go to https://huggingface.co/kyutai/pocket-tts and accept the terms, then make sure you're logged in locally with `uvx hf auth login`.c                       e Zd ZdZdZ	 	 	 	 d>dededededz  d	ed
e	dz  de
dedz  de
f fdZedej                  fd       Zedefd       Zed	ededz  d
e	dz  defd       Ze	 d?d	ededz  d
e	dz  defd       Zeddeeeedfdedz  d	ee	z  dz  deez  dedeez  dz  dede
defd       Z	 	 	 d@dedej6                  dz  dej6                  dz  dej6                  dz  deej6                  ej6                  f   f
dZdedej6                  dej6                  dej6                  deej6                  ej6                  f   f
dZdej6                  defd Zd!ej6                  dej6                  fd"Z ded#eddfd$Z!dedefd%Z"ejF                  d&e$jJ                  d'e$jJ                  d(ed)efd*       Z&ejF                  e'dd+fded,ed-ed.edz  d/e
dej6                  fd0       Z(ejF                  e'dd+fded,ed-ed.edz  d/e
f
d1       Z)ejF                  ded,ed.ed/e
fd2       Z*ejF                  ded3e+d4ed.ed&e$jJ                  d'e$jJ                  fd5       Z,ejF                  ded4ed.ed&e$jJ                  fd6       Z- e.d78      	 dAde	ez  ej6                  z  d9e
defd:       Z/ejF                  	 dAde	ez  ej6                  z  d9e
defd;       Z0d<edefd=Z1 xZ2S )BTTSModelg      @g       @NFflow_lmtemplsd_decode_stepsnoise_clampconfigorigin pad_with_spaces_for_short_inputs"model_recommended_frames_after_eosremove_semicolonsc                     t         |           || _        || _        || _        || _        || _        || _        d| _        || _	        || _
        |	| _        |
| _        y )NT)super__init__r*   r+   r,   r-   eos_thresholdr.   has_voice_cloningr/   r0   r1   r2   )selfr*   r+   r,   r-   r6   r.   r/   r0   r1   r2   	__class__s              l/Users/ahmed/devFolder/Ultron/claude-voice/.venv/lib/python3.12/site-packages/pocket_tts/models/tts_model.pyr5   zTTSModel.__init__B   sd     		 0&*!%6V-2T/!2    returnc                 H    t        | j                               j                  S N)next
parametersdevicer8   s    r:   rA   zTTSModel.device\   s    DOO%&---r;   c                 B    | j                   j                  j                  S r>   )r.   mimisample_raterB   s    r:   rE   zTTSModel.sample_rate`   s    {{+++r;   c                    t        j                  |j                  |j                  j                  j
                  |j                  j                        } | ||||||||j                  |j                  |j                  
      }|S )N)
latent_diminsert_bos_before_voice)r/   r0   r1   r2   )
r   from_pydantic_configr*   rD   	quantizer	dimensionrH   r0   r1   r2   )	clsr.   r+   r,   r-   r6   r/   r*   	tts_models	            r:   _from_pydantic_configzTTSModel._from_pydantic_configd   s     22NN{{,,66$*NN$J$J

 -3-T-T/5/X/X$66
	 r;   c                 ^	   | j                  ||||||      }t        j                  j                  t        j                  |j
                  j                  j                  |j                  j                  xs  |j                  j                  j                  ft        j                              |j
                  _        |j
                  j                  |j                  j                  t        d      t         j#                  d|j
                  j                          t%        t'        |j
                  j                              }|j
                  j)                  |d       |j                  j+                         }	t-        di |	d   }
t/        di |	d   }t1        j2                  di |	d   }t1        j2                  di |	d   }t5        di |	d	   }t7        |
|||	d
   |	d   |	d   |	d   |
j8                  z  |	d   |	d   ||      j;                  d      |_        |j                  j                  |j
                  j                  t        d      t         j#                  d|j                  j                          t=        t'        |j                  j                              }|j                  j)                  |d       |j                  j?                          |j                  jt         j#                  d|j                          	 t'        |j                        }tF        j                  jI                  |      }|j)                  |d       |j
                  j                  !|j                  t         jK                  d       tM        |jO                               dz  }tP        jR                  jU                  dd      dk(  rHd}tF        j                  jW                  |jO                         |       t         j#                  d|        tY        j"                  d| d       |j
                  |j                  fD ]2  }|j[                         D ]  \  }}t]        |t^              s||_0         4 |S # t@        $ r  d|_!        t'        |jD                        }Y w xY w) Nr/   )dtypezHIf you specify flow_lm.weights_path you should specify mimi.weights_pathzLoading FlowLM weights from T)strictseanettransformerrJ   channelsrE   
frame_rate	inner_dim	outer_dim)rU   rE   rV   encoder_frame_raterW   rX   encoder_transformerdecoder_transformercpu)rA   zHIf you specify mimi.weights_path you should specify flow_lm.weights_pathzLoading Mimi weights from zLoading TTSModel weights from FzINo weights_path specified for FlowLM or TTSModel, model is uninitialized!    .APOCKET_TTS_SAVE_WEIGHTS01z./model.safetensorszSaved TTSModel weights to z+TTS Model loaded successfully. Its size is z MB )1rN   torchr   	Parameterzerosr*   rT   d_modelrD   rW   rS   rK   float32speaker_proj_weightweights_path
ValueErrorloggerinfor%   r"   load_state_dict
model_dumpr   r   r   ProjectedTransformerr   r   
hop_lengthtor&   eval	Exceptionr7   "weights_path_without_voice_cloningsafetensors	load_filewarningr$   
state_dictosenvironget	save_fileloggingnamed_modules
isinstancer   _module_absolute_name)rL   r.   r+   r,   r-   r6   r/   rM   state_dict_flowlmmimi_configencoderdecoderrZ   r[   rJ   
mimi_stateweights_filerw   
size_in_mb	save_path
top_modulemodule_namemodules                          r:   "_from_pydantic_config_with_weightsz+TTSModel._from_pydantic_config_with_weights   s*    --D*Kv . 
	 160B0BKKNN..66KK))IV[[-?-?-I-I mm1
	- >>&&2{{''/ ^  KK6v~~7R7R6STU 6%fnn&A&AB! --.?-M kk,,.  8+h"788+h"78.CCakR_F`a.CCakR_F`a">[%=>	" ,#M2"<0*=9G<N<NN!+.!+. 3 3
 "E"
 	  ;;##/~~**2 ^  KK4V[[5M5M4NOP,-B6;;C[C[-\]JNN**:d*C*KK89L9L8MNO`4V5H5HI
 %**44\BJ%%j%>>>&&.63F3F3NNN[ ")"6"6"89S@
::>>3S9S@-I''	(<(<(>	JKK4YK@AB:,cRS %,,inn=J'1'?'?'A#V!&.9/:, (B > 3  `.3	+4V5^5^_`s   R %R,+R,languager6   quantizec                    ||t        d      ||t        }||dk(  rt        d      t        | dz  }t        |      }|j                  dvrt        d      t        |      }t        |      }t        j                  d| d       t        j                  ||||||	      }	|rt        |	j                  t               |	S )
a	  Load a pre-trained TTS model with specified configuration.

        This class method loads a complete TTS model including the flow language model
        and Mimi compression model from pre-trained weights. The model is initialized
        with the specified generation parameters and ready for inference.

        Args:
            language: Optional language identifier to select a predefined config. Incompatible with
                the `config` argument. Available options
                are `"english_2026-01"`, `"english_2026-04"`, `"english"`, `"french_24l"`, `"german_24l"`, `"portuguese"`, `"italian"`, `"spanish_24l"`.
                If neither `config` nor `language` is provided, defaults to `"english", which is the same model as 'english_2026-04'`.
            config: A path to a custom YAML config file saved locally (e.g., `"C://pocket_tts/pocket_tts_config.yaml"`).
            temp: Sampling temperature for generation. Higher values produce more
                diverse but potentially lower quality output.
            lsd_decode_steps: Number of steps for Lagrangian Self Distillation
                decoding. More steps can improve quality but increase computation.
            noise_clamp: Maximum value for noise sampling. If None, no clamping
                is applied. Helps prevent extreme values in generation.
            eos_threshold: Threshold for end-of-sequence detection. Higher values
                make the model more likely to continue generating.
            quantize: If True, apply dynamic int8 quantization to the transformer's
                attention and FFN layers. Reduces runtime memory by ~48% and improves
                inference speed by ~27% on x86 (FBGEMM).
                No measurable impact on speech quality (WER unchanged).
                For optimized performance, install torchao: ``pip install pocket-tts[quantize]``

        Returns:
            TTSModel: Fully initialized model with loaded weights on cpu, ready for
                text-to-speech generation.

        Raises:
            FileNotFoundError: If the specified config file or model weights
                are not found.
            ValueError: If the configuration is invalid or incompatible.

        Example:
            ```python
            from pocket_tts import TTSModel

            # Load with default settings
            model = TTSModel.load_model()

            # Load with int8 quantization
            model = TTSModel.load_model(quantize=True)
            ```
        zHCannot specify both config and language, please choose one or the other.frenchzzFor technical reasons, only a larger 24-layer model is available for French. Please use the 'french_24l' language instead..yaml)r   z.ymlz8Config should be a path to a YAML file ending with .yamlzLoading model from config at z...rP   )ri   r   r   r   suffixr   rj   rk   r)   r   r   r*   r   )
rL   r   r.   r+   r,   r-   r6   r   config_pathrM   s
             r:   
load_modelzTTSModel.load_model   s    r ("6Z  >h.'H8#  Q  !hZu#55Ff== 11WXX6l[)3K=DE??D*K{ @ 
	 y002DEr;   model_statetext_tokensbackbone_input_latentsaudio_conditioningc                    |:t        j                  dt         j                  | j                  j                        }|Wt        j
                  dd| j                  j                  f| j                  j                  | j                  j                        }|Wt        j
                  dd| j                  j                  f| j                  j                  | j                  j                        }| j                  ||||      }|j                  d   |j                  d   z   |j                  d   z   }t        | j                  ||       |S )zJFirst one is the backbone output, second one is the audio decoding output.)r'   r   )rQ   rA   r'   r   )r   r   r   r   	increment)rb   rd   int64r*   rA   emptyldimrQ   dim_run_flow_lmshaper   )r8   r   r   r   r   outputincrement_bys          r:   _run_flow_lm_and_increment_stepz(TTSModel._run_flow_lm_and_increment_step=  s     ++fEKKH[H[\K!)%*[[At||(()1C1CDLLL_L_&" %!&At||''(0B0B4<<K^K^" ""##9#1	 # 
 a #9#?#?#BBEWE]E]^_E`` 	 	k\Jr;   c           	      2   | j                   j                  t        |            }t        j                  ||gd      }| j                   j                  |||| j                  | j                  | j                  | j                        \  }}|d d d d d f   |fS )Nr'   r   )r   r,   r+   r-   r6   )
r*   conditionerr   rb   cat_sample_next_latentr,   r+   r-   r6   )r8   r   r   r   r   text_embeddingsoutput_embeddingsis_eoss           r:   r   zTTSModel._run_flow_lm\  s     ,,22=3MN))_6H$IqQ$(LL$D$D"#!22((,, %E %
!6 !D!,f44r;   encodedfilenamec                    t        | j                  dd      }|j                  d   | j                  j                  j                  k(  r| j                  j                  |      }n|}| j                  j                  ||      }t        j                  j                  j                  || j                  |j                                t        j                  d|       y )Nr'   i'  
batch_sizesequence_lengthz;Saved restored audio from Mimi encoding to %s for debugging)r   rD   r   rJ   rK   decode_from_latentscipyiowavfilewriterE   numpyrj   rk   )r8   r   r   r   latent_to_decoderesored_audios         r:   _decode_and_dumpzTTSModel._decode_and_dumpq  s     q%P
==tyy22<<<#yy227;&		445EzRx)9)9=;N;N;PQQS[\r;   audioc                    | j                   j                  |      }t        r| j                  |d       |j	                  dd      j                  t        j                        }t        j                  || j                  j                        }|S )Nz debug_encoded_latent_decoded.wav)rD   encode_to_latentr    r   	transposerp   rb   rf   Flinearr*   rg   )r8   r   r   latentsconditionings        r:   _encode_audiozTTSModel._encode_audio{  sk    )),,U3!!'+MN##B+..u}}=xx)I)IJr;   r   c           	         |j                         D ]  \  }}d|v s|d   }|j                  d   }||k  s%t        j                  |j                  d   |j                  d   ||j                  d   |j                  d   ft	        d      |j
                  |j                        }||d	d	d	d	d	|d	d	d	d	f<   ||d<    y	)
a  Expand KV cache back to full sequence_length for generation.

        When a model state is retrieved from cache with sliced KV caches,
        this method expands them back to the full size needed for generation.

        Args:
            model_state: The model state dict containing potentially sliced KV caches
            sequence_length: Target sequence length to expand caches to
        cache   r   r'         NaN)rA   rQ   N)itemsr   rb   fullfloatrA   rQ   )r8   r   r   r   module_stater   current_lengthexpanded_caches           r:   _expand_kv_cachezTTSModel._expand_kv_cache  s     *5):):)<%K,&$W-!&Q!O3%*ZZ!KKN!KKN+!KKN!KKN e$||#kk&N CHN1a.!Q#>?,:L)+ *=r;   c                     |j                         D ]B  }|j                  d      }|t        |j                  d      d   j	                               c S  t        d      )Noffsetr   r   znCould not find offset in model state, please open an issue at https://github.com/kyutai-labs/pocket-tts/issues)valuesrz   intviewitemri   )r8   r   r   r   s       r:   _flow_lm_current_endzTTSModel._flow_lm_current_end  sa    '..0L!%%h/F!6;;r?1-22455 1 B
 	
r;   latents_queueresult_queuemimi_sequence_lengthmimi_steps_per_latentc                 v   	 g }t        | j                  d|      }	 |j                         }|nO|| j                  j                  z  | j                  j
                  z   }|j                  dd      }	| j                  j                  |	      }
t        j                         }| j                  j                  |
|      }t        | j                  ||       |j                  d   | j                  j                  j                  z  }t        j!                  dt#        |d	z        t#        t        j                         |z
  d	z               |j%                  |       |j'                  d
|f       |j)                          c|j'                  d       y# t*        $ r}|j'                  d|f       Y d}~yd}~ww xY w)zVWorker thread function for decoding audio latents from queue with immediate streaming.r'   r   Nr   r   r   r   zG                              Decoded %d ms of audio with mimi in %d ms  chunk)doneNerror)r   rD   rz   r*   emb_stdemb_meanr   rJ   time	monotonicr   r   r   r.   rE   rj   debugr   appendput	task_donerr   )r8   r   r   r   r   audio_chunksr   latentmimi_decoding_input
transposed	quantizedtaudio_frameaudio_frame_durationes                  r:   _decode_audio_workerzTTSModel._decode_audio_worker  sz    	+L$TYY1NbcJ&**,>&,t||/C/C&CdllF[F[&[#0::2rB
 II//
;	NN$"ii::9jQ		:AVW'2'8'8';dkk>N>N>Z>Z'Z$J,t34)A-56
 ##K0  ';!78'')- 2 ^, 	+gq\**	+s   FF 	F8F33F8Ttext_to_generate
max_tokensframes_after_eos
copy_statec                     g }| j                  |||||      D ]  }|j                  |        t        j                  |d      S )a  Generate complete audio tensor from text input.

        This method generates the full audio output for the given text prompt
        and returns it as a single tensor. It internally uses the streaming
        generation method but collects all chunks before returning.

        This method is NOT thread-safe; separate model instances should be used
        for concurrent generation.

        Args:
            model_state: Model state dictionary containing hidden states and
                positional information. Can be obtained from get_state_for_audio_prompt()
                or init_states(). The state may be modified during generation.
            text_to_generate: Input text to convert to speech. The text will be
                automatically formatted (capitalization, punctuation) for optimal
                generation quality.
            frames_after_eos: Number of additional frames to generate after
                detecting end-of-sequence. If None, automatically determined
                based on text length (1-3 frames).
            copy_state: Whether to create a deep copy of the model state before
                generation. If True, preserves the original state for reuse.
                If False, modifies the input state in-place. Defaults to True.

        Returns:
            torch.Tensor: Generated audio tensor with shape [channels, samples]
                at the model's sample rate (typically 24kHz). The audio is
                normalized and ready for playback or saving.
                You can get the sample rate from the `sample_rate` attribute.

        Raises:
            ValueError: If text_to_generate is empty or invalid.
            RuntimeError: If generation fails due to model errors.

        Example:
            ```python
            from pocket_tts import TTSModel

            model = TTSModel.load_model()

            voice_state = model.get_state_for_audio_prompt("hf://kyutai/tts-voices/alba-mackenna/casual.wav")

            # Generate audio
            audio = model.generate_audio(voice_state, "Hello world!", frames_after_eos=2, copy_state=True)

            print(f"Generated audio shape: {audio.shape}")
            print(f"Audio duration: {audio.shape[-1] / model.sample_rate:.2f} seconds")
            ```
        )r   r   r   r   r   r   r   )generate_audio_streamr   rb   r   )r8   r   r   r   r   r   r   r   s           r:   generate_audiozTTSModel.generate_audio  sV    r //#--!! 0 
E &
 yy1--r;   c              #   V  K   || j                   }t        | j                  j                  j                  ||| j
                  | j                        }|D ]N  }t        || j
                  | j                        \  }}|dz  }||n|}	| j                  |||	|      E d{    P y7 w)a	  Generate audio streaming chunks from text input.

        This method generates audio from text and yields chunks as they become
        available, enabling real-time playback or processing. It uses multithreading
        to parallelize generation and decoding for optimal performance.
        This method is NOT thread-safe; separate model instances should be used
        for concurrent generation.

        Args:
            model_state: Model state dictionary containing hidden states and
                positional information. Can be obtained from get_state_for_audio_prompt()
                or init_states(). The state may be modified during generation.
            text_to_generate: Input text to convert to speech. The text will be
                automatically formatted (capitalization, punctuation) for optimal
                generation quality.
            frames_after_eos: Number of additional frames to generate after
                detecting end-of-sequence. If None, automatically determined
                based on text length (1-3 frames). Defaults to None.
            copy_state: Whether to create a deep copy of the model state before
                generation. If True, preserves the original state for reuse.
                If False, modifies the input state in-place. Defaults to True.

        Yields:
            torch.Tensor: Audio chunks with shape [samples] at the model's
                sample rate (typically 24kHz). Chunks are yielded as soon as
                they are decoded, enabling real-time streaming.

        Raises:
            ValueError: If text_to_generate is empty or invalid.
            RuntimeError: If generation fails due to model errors or threading issues.

        Example:
            ```python
            from pocket_tts import TTSModel

            model = TTSModel.load_model()

            voice_state = model.get_state_for_audio_prompt("hf://kyutai/tts-voices/alba-mackenna/casual.wav")
            # Stream generation
            for chunk in model.generate_audio_stream(voice_state, "Long text content..."):
                # Process each chunk as it's generated
                print(f"Generated chunk: {chunk.shape[0]} samples")
                # Could save chunks to file or play in real-time
            ```

        Note:
            This method uses multithreading to parallelize latent generation
            and audio decoding. Generation performance is logged including
            real-time factor (RTF) metrics.
        N)r2   r   )r   r   r   r   )	r1   split_into_best_sentencesr*   r   	tokenizerr0   r2   prepare_text_prompt!_generate_audio_stream_short_text)
r8   r   r   r   r   r   chunksr   frames_after_eos_guesseffective_framess
             r:   r   zTTSModel.generate_audio_stream   s     v ##FF +LL$$..11"44
 E7Jt<<d>T>T844 #a'"$4$@ F\  =='!&!1%	 >    s   BB)B' B)c              #     K   |rt        j                  |      }| j                  j                  j	                  |      }|j
                  j                  d   }| j                  |      }t        | j                  j                  | j                  j                  z        }||z  }	t        j                         }
t        j                         }t        j                  | j                   |
||	|fd      }t"        j%                  d       t'        j(                         }|j+                          | j-                  |||||
|       d}	 |j/                         }|d   dk(  r|d   }||j                  d   z  }|d	    n:|d   d
k(  rn2|d   dk(  r)t1        d      5  |j3                          d d d        |d   rt1        d      5  |j3                          d d d        t        |dz  | j4                  j                  j6                  z        }t        t'        j(                         |z
  dz        }||z  }t"        j%                  d|||       y # 1 sw Y   |d   xY w# 1 sw Y   xY ww)Nr'   T)targetargsdaemonzstarting timer now!)r   preparedmax_gen_lenr   r   r   r   r   r   )r   r   r   r   z"Waiting for mimi decoder to finishr   zAGenerated: %d ms of audio in %d ms so %.2fx faster than real-time)copydeepcopyr*   r   preparetokensr   _estimate_max_gen_lenr   rD   rY   rV   queueQueue	threadingThreadr   rj   rk   r   r   start	_generaterz   r!   joinr.   rE   )r8   r   r   r   r   r  token_countr	  r   r   r   r   decoder_threadt_generatingtotal_generated_samplesresultaudio_chunkduration_generated_audiogeneration_timereal_time_factors                       r:   r   z*TTSModel._generate_audio_stream_short_texty  sO     --4K<<++334DEoo++A.00= #DII$@$@499CWCW$W X*-BB {{} #)),,/CEZ[

 	)*~~' 	##-'% 	 	
 #$!%%'FayG#$Qi';+<+<R+@@'!$''f$g%+,PQ"'') R Qi $ $$HI! J $'#d*T[[-=-=-I-II$
  t~~/,>$FG3oEO$		
! R Qi JIs7   FI'I)I'IA8I'II'I$ I'r  r	  c                 v    |j                   j                  d   } j                        }||z   z   }	 j                  |	       t	        d      5   j                  |j                          d d d         fd}
t        j                  |
d      }|j                          y # 1 sw Y   ;xY w)Nr'   )r   zPrompting text)r   r   c                      	 j                         y # t        $ rO} t        j                  d|         j	                  d        j	                  d| f       Y d } ~ y Y d } ~ y d } ~ ww xY w)Nz$Error in autoregressive generation: r   )_autoregressive_generationrr   rj   r   r   )r   r   r   r	  r   r   r8   s    r:   run_generationz*TTSModel._generate.<locals>.run_generation  sz    3//.>  3CA3GH ,!%%d++ $$gq\22 ,3s    	A0A A++A0T)r  r  )	r  r   r   r   r!   r   r  r  r  )r8   r   r  r	  r   r   r   r  current_endrequired_lenr"  generation_threads   `` ````     r:   r  zTTSModel._generate  s     oo++A.//<"[0;>k<H#$4500'X__ 1  6
	3 	3 &,,N4P!) 65s   B//B8c           
      X   t        j                  dd| j                  j                  ft	        d      t        t        | j                  j                                     j                  | j                  j                        }g }d }t        |      D ]  }t        dd      5 }	| j                  ||      \  }
}|j                         r||}||||z   k\  r	 d d d         n{|j                  |
       |
}d d d        |j                  	j                           t"        j$                  j'                  dd	      d
k(  rt)        d      t*        j-                  d       |j                  d        t*        j/                  dt1        t3        j4                  |                   y # 1 sw Y   xY w)Nr'   r   )
fill_valuerA   rQ   zGenerating latentF)print_output)r   r   KPOCKET_TTS_ERROR_WITHOUT_EOSr_   r`   z.Generation reached maximum length without EOS!zRMaximum generation length reached without EOS, this very often indicates an error.z#Average generation step time: %d ms)rb   r   r*   r   r   r?   iterr@   rA   rQ   ranger!   r   r   r   r   elapsed_time_msrx   ry   rz   RuntimeErrorrj   rv   rk   r   
statisticsmean)r8   r   r	  r   r   backbone_inputsteps_timeseos_stepgeneration_steptimernext_latentr   s               r:   r!  z#TTSModel._autoregressive_generation  st    4<<$$%U|T\\44678??,,$$	
 $[1O'(;%PTY&*&J&J +N 'K '#V ;;=X%5.H'OxJZ?Z,Z QP !!+.!, Q u445  2 zz~~=sCsJ"#STTNNd
 	$93z{?[;\]- QPs   !6F !F  F)	r   )maxsizetruncatec                 &    | j                  ||      S r>   )get_state_for_audio_prompt)r8   r   r7  s      r:   "_cached_get_state_for_audio_promptz+TTSModel._cached_get_state_for_audio_prompt  s     ../A8LLr;   c                    t        |t        t        f      rKt        |      j                  d      r1t        |t              rt	        |      }t        || j                        S t        |t              r|t        v r| j                  | j                  j                  t              st        d| j                         t        t	        t        | j                  j                  |            | j                        S | j                  s%t        |t        t        f      rt        t              t        |t              rt	        |      }t        |t              r~t!        |      \  }}|rBt#        d|z        }|j$                  d   |kD  r"|dd|f   }t&        j)                  d| d	       t+        ||| j,                  j.                  j0                  d
      }t3        d      5  | j5                  |j7                  d      j9                  | j                              }ddd       | j:                  j<                  r-t?        j@                  | j:                  jB                  gd
      }tE        | j:                  d
j$                  d
         }t3        d      5  | jG                  ||       ddd       t&        j)                  dtI        |      dz         |S # 1 sw Y   xY w# 1 sw Y   9xY w)a 
  Create model state conditioned on audio prompt for continuation.

        This method processes an audio prompt and creates a model state that
        captures the acoustic characteristics (speaker voice, style, prosody)
        for use in subsequent text-to-speech generation. The resulting state
        enables voice cloning and audio continuation with speaker consistency.

        Args:
            audio_conditioning: Audio prompt to condition (or .safetensors to load). Can be:
                - Path: Local file path to audio file (or .safetensors)
                - str: URL to download audio file (or .safetensors) from
                - torch.Tensor: Pre-loaded audio tensor with shape [channels, samples]
            truncate: Whether to truncate long audio prompts to 30 seconds.
                Helps prevent memory issues with very long inputs. Defaults to False.

        Returns:
            dict: Model state dictionary containing hidden states and positional
                information conditioned on the audio prompt. This state can be
                passed to `generate_audio()` or `generate_audio_stream()` for
                voice-consistent generation.

        Raises:
            FileNotFoundError: If audio file path doesn't exist.
            ValueError: If audio tensor is invalid or empty.
            RuntimeError: If audio processing or encoding fails.

        Example:
            ```python
            from pocket_tts import TTSModel

            model = TTSModel.load_model()
            # From HuggingFace URL
            voice_state = model.get_state_for_audio_prompt("hf://kyutai/tts-voices/alba-mackenna/casual.wav")

            # From local file
            voice_state = model.get_state_for_audio_prompt("./my_voice.wav")

            # Reload state from a .safetensors file (much faster than extracting from an audio file)
            voice_state = model.get_state_for_audio_prompt("./my_voices.safetensors")

            # From HTTP URL
            voice_state = model.get_state_for_audio_prompt(
                "https://huggingface.co/kyutai/tts-voices/resolve"
                "/main/expresso/ex01-ex02_default_001_channel1_168s.wav"
            )
            ```

        Note:
            - Audio is automatically resampled to the model's sample rate (24kHz)
            - The audio is encoded using the Mimi compression model and projected
              to the flow model's latent space
            - Processing time is logged for performance monitoring
            - The state preserves speaker characteristics for voice cloning
        z.safetensorsNzvCannot use predefined voices when the model is not loaded from a config associated with a language.Here the origin is )r   name   r   .z%Audio truncated to first 30 seconds (z	 samples)r'   zEncoding audio promptr   r   r   zPrompting audio)r   r   z/Size of the model state for audio prompt: %d MBr]   )%r~   strr   endswithr"   _import_model_staterA   r   r/   is_relative_tor   ri   r#   stemr7   VOICE_CLONING_UNSUPPORTEDr	   r   r   rj   rk   r
   r.   rD   rE   r!   r   	unsqueezerp   r*   rH   rb   r   bos_before_voicer   r   r$   )r8   r   r7  r   conditioning_sample_ratemax_samplespromptr   s           r:   r9  z#TTSModel.get_state_for_audio_prompt  s   t (3+63?Q;R;[;[<
 ,c2%:;M%N"&'94;;GG )3/"&CC {{"$++*D*D[*Q **.++8 
 '%($++2B2BI[\ 	  %%*5G#t*U677(#.!67I!J($/.89K.L+E+!"'?"?@;;r?[0!#||"34EKK"G}T] ^_!./1A1A1M1Mq" $$;<''(:(D(DQ(G(J(J4;;(WXF = <<//YY = =vFANF!$,,1fll[\o^#$5600[]c0d 7 	=|K?X\_?_	
  =< 76s    :KK!K!K*r  c                     || j                   z  | j                  z   }| j                  j                  j                  }t        j                  ||z        S r>   )_TOKENS_PER_SECOND_ESTIMATE_GEN_SECONDS_PADDINGr.   rD   rV   mathceil)r8   r  gen_len_secrV   s       r:   r  zTTSModel._estimate_max_gen_len  sF    !D$D$DDtG`G``[[%%00
yyz122r;   )NFNFr>   )NNN)F)3__name__
__module____qualname__rJ  rK  r   r   r   r   r   boolr5   propertyrb   rA   rE   classmethodr   rN   r   r   r   r   r   r>  r   dictTensortupler   r   r   r   r   r   no_gradr  r  r   r   r   r   r   r   r  r!  r   r:  r9  r  __classcell__)r9   s   @r:   r)   r)   >   s   "% #169="'33 3 	3
 T\3 3 t3 +/3 -0$J3  34 . . . ,S , , 
 T\ t 
 8  #dd
 T\d td 
d dL   $$(/ 8*=4R*R d
T!R ck	R
 R S[4'R R R 
R Rn ,06:26 \\D( !&t 3	
 "LL4/ 
u||U\\)	*>55 \\5 !&	5
 "LL5 
u||U\\)	*5*] ] ]	5<< 	ELL 	;D ;3 ;4 ;B
 
 
 ]](+{{(+ kk(+ "	(+
  #(+ (+T ]]
 .'+A.A. A. 	A.
 *A. A. 
A. A.F ]]
 .'+VV V 	V
 *V V Vp ]]G
G
36G
JMG
[_G
 G
R ]]""""  "" 	""
 "" {{"" kk"" ""H ]]"^"^.1"^EH"^Y^YdYd"^ "^H qNSM"&*u||";MGKM	M M
 ]]NSu"&*u||";uGKu	u un3 3 3r;   r)   textr0   r2   r<   c                    | j                         } | dk(  rt        d      | j                  dd      j                  dd      j                  dd      } |r| j                  dd      } t        | j	                               }|d	k  rd
}nd}| d   j                         s| d   j                         | dd  z   } | d   j                         r| dz   } |r!t        | j	                               dk  rd| z   } | |fS )N zText prompt cannot be empty
 z  ;,r   r   r'   r   r   .   z        )stripri   replacelensplitisupperupperisalnum)rZ  r0   r2   number_of_wordsr  s        r:   r   r     s     ::<Drz677<<c"**45==dCHD||C%$**,'O!!"!" 7??Aw}}ab) Bxcz (C

,=,A~'''r;   list_of_tokensboundary_tokensc                     dg}d}t        |       D ]!  \  }}||v rd}|r|j                  |       d}# |j                  t        |              |S )a&  Find token indices where text should be split based on boundary tokens.

    Returns a list of boundary positions used to slice segments. Each consecutive
    pair (indices[i], indices[i+1]) defines one segment. The first element is
    always 0 and the last is always len(list_of_tokens).
    r   FT)	enumerater   rf  )rl  rm  indicesprevious_was_boundaryidxtokens         r:   _find_boundary_indicesrt    s_     cG!/
UO#$(!$s#$)! 0 NN3~&'Nr;   boundary_indicesc                     g }t        t        |      dz
        D ]C  }||   }||dz      }|j                  j                  | ||       }|j	                  ||z
  |f       E |S )zNDecode token segments between boundary indices into (token_count, text) pairs.r'   )r+  rf  spdecoder   )rl  ru  r   segmentsir  endrZ  s           r:   _segments_from_boundariesr|    sr     H3'(1,- #q1u%||"">%#<=ud+,	 .
 Or;   r   r   c                 \   t        |||      \  }}|j                         } | |      }|j                  d   j                         } | d      j                  d   j                         ^}}t	        ||      }	t        ||	|       }
 | d      j                  d   j                         ^}}g }|
D ]  \  }}||k  r|j                  ||f        | |j                               j                  d   j                         }t	        ||      }t        |||       }t        |      dkD  r|j                  |       |j                  ||f        |}g }d}d}|D ]H  \  }}|dk(  r|}|}||z   |kD  r$|j                  |j                                |}|}<|d|z   z  }||z  }J |dk7  r|j                  |j                                |D ]c  } | |j                               j                  d   j                         }t        |      |kD  sCt        j                  dt        |      ||       e |S )Nr   z.!...?z,;:r'   r\  r^  zCChunk has %d tokens (max %d), generation may skip words: '%.50s...')r   rd  r  tolistrt  r|  r   rf  extendrj   rv   )r   r   r   r0   r2   _r  rl  end_of_sentence_tokenssentence_boundariesnb_tokens_and_sentencesfallback_tokensrefined_segments	nb_tokensrZ  
sub_tokenssub_boundariessub_segmentsmax_nb_tokens_in_a_chunkr  current_chunkcurrent_nb_of_tokens_in_chunksentencer   chunk_tokenss                            r:   r   r     sX    .:<Ma (--/'(F]]1%,,.N!*8!4!;!;A!>!E!E!GA0AWX7+Y
 $E*11!4;;=A2	4
"##Y$56"4::<077:AACJ3JPN4ZQZ[L< 1$ ''5 ''D(9: 3  *FM$%!/	8B$M,5)(947OOMM---/0$M,5)S8^+M)Y6)  0 m))+, /66q9@@B|z)NNUL!	  Mr;   r   destc                     i }| j                         D ]'  \  }}|j                         D ]  \  }}||| d| <    ) t        j                  j                  ||       y )N/)r   rt   rb   r{   )r   r  dict_to_storer   r   keytensor_values          r:   export_model_stater    sd    M%0%6%6%8!\!-!3!3!5C4@M[M3%01 "6 &9 t4r;   sourcerA   c                    i }t        j                  | d      5 }|j                         D ]  }|j                  d      \  }}|j	                  |i        |dk(  rL|j                  |      }t        j                  d|j                  d   t        j                  |      ||   d<   z|j                  |      j                  |      ||   |<    	 d d d        |S # 1 sw Y   |S xY w)	Npt)	frameworkr  r#  )r'   r   )r'  rQ   rA   r   )rt   	safe_openkeysrg  
setdefault
get_tensorrb   r   r   longrp   )r  rA   r  fr  r   
tensor_keytensors           r:   r@  r@    s     F			v	6!668C&)iin#Kk2.]* c*05

V\\!_EJJv1{#H- 34,,s2C2F2Fv2N{#J/  
7 M 
7 Ms   B3CC")Zr
  r|   rL  rx   r  r.  r  r   	functoolsr   pathlibr   rt   safetensors.torchscipy.io.wavfiler   rb   r   torch.nnr   r   typing_extensionsr   pocket_tts.conditioners.baser   pocket_tts.data.audior	   pocket_tts.data.audio_utilsr
   pocket_tts.default_parametersr   r   r   r   r   r   pocket_tts.models.flow_lmr   pocket_tts.models.mimir   pocket_tts.modulesr   "pocket_tts.modules.dummy_quantizerr   pocket_tts.modules.seanetr   r   "pocket_tts.modules.stateful_moduler   r   r   pocket_tts.quantizationr   r   pocket_tts.utils.configr   r   r   pocket_tts.utils.utilsr   r    r!   r"   r#   r$    pocket_tts.utils.weights_loadingr%   r&   set_num_threads	getLoggerrO  rj   listr  rC  Moduler)   r>  rR  rW  r   r   rt  r|  r   rU  rV  r  rA   r@  ra   r;   r:   <module>r     s      	            $ " 6 , 5  2 , / = B [ [ J D D  Y   a 			8	$@@DEgEbEgEgEi@j?k lHI P3ryy P3f(
(15(JN(
38_(@49 tCy UYZ]U^ (
I
15c
	%S/
BB B '+	B
 B 
#YBJ5Dd33D.E)E$F 5cTXj 5$J %	#tC%&
&'r;   