Ë
    ÿÍ:jµ5  ã                   óv  — d dl Z d dlZd dlmZmZ d dlmZ d dlmZ d dl	m
Z
mZmZ d dlZd dlZd dlmZ d dlmZmZmZ g Zd ej,                   ej.                  ej0                  «      j2                  «      z  Z ed	d
ez  «      Zd„ Z G d„ dej<                  j>                  «      Z  G d„ dej<                  j>                  «      Z! G d„ de«      Z" G d„ de«      Z# G d„ dej<                  j>                  e"«      Z$ G d„ de#«      Z%e G d„ d«      «       Z& e&d eed¬«      dddd d!d"d#d$d%d¬&«      Z'd'e'_(        y)(é    N)ÚABCÚabstractmethod)Ú	dataclass)Úpartial)ÚCallableÚListÚTuple)Úmodule_utils)Úemformer_rnnt_baseÚRNNTÚRNNTBeamSearché(   é
   gš™™™™™©?c                 óö   — t        j                  | | t        j                  kD     «      | | t        j                  kD  <   | | t        j                  k     t        j                  z  | | t        j                  k  <   | S ©N)ÚtorchÚlogÚmathÚe©Úxs    úw/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchaudio/pipelines/rnnt_pipeline.pyÚ_piecewise_linear_logr      sS   € Ü—I‘I˜a ¤D§F¡F¡
™mÓ,€A€aŒ$�&‰&�j�MØ�qœDŸF™F‘{‘^¤d§f¡fÑ,€A€aŒ4�6‰6�k�NØ€Hó    c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )Ú_FunctionalModulec                 ó0   •— t         ‰| �  «        || _        y r   )ÚsuperÚ__init__Ú
functional)Úselfr    Ú	__class__s     €r   r   z_FunctionalModule.__init__   s   ø€ Ü‰ÑÔØ$ˆ�r   c                 ó$   — | j                  |«      S r   )r    ©r!   Úinputs     r   Úforwardz_FunctionalModule.forward   s   € Ø�‰˜uÓ%Ð%r   ©Ú__name__Ú
__module__Ú__qualname__r   r&   Ú__classcell__©r"   s   @r   r   r      s   ø„ ô%ö&r   r   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )Ú_GlobalStatsNormalizationc                 óH  •— t         ‰| �  «        t        |«      5 }t        j                  |j                  «       «      }d d d «       | j                  dt        j                  d   «      «       | j                  dt        j                  |d   «      «       y # 1 sw Y   ŒZxY w)NÚmeanÚ	invstddev)	r   r   ÚopenÚjsonÚloadsÚreadÚregister_bufferr   Útensor)r!   Úglobal_stats_pathÚfÚblobr"   s       €r   r   z"_GlobalStatsNormalization.__init__$   s   ø€ Ü‰ÑÔäÐ#Ó$ð 	(¨Ü—:‘:˜aŸf™f›hÓ'ˆD÷	(ð 	×Ñ˜V¤U§\¡\°$°v±,Ó%?Ô@Ø×Ñ˜[¬%¯,©,°t¸KÑ7HÓ*IÕJ÷		(ð 	(ús   ›$BÂB!c                 ó:   — || j                   z
  | j                  z  S r   )r0   r1   r$   s     r   r&   z!_GlobalStatsNormalization.forward-   s   € Ø˜Ÿ	™	Ñ! T§^¡^Ñ3Ð3r   r'   r,   s   @r   r.   r.   #   s   ø„ ôKö4r   r.   c                   ól   — e Zd Zedej
                  deej
                  ej
                  f   fd„«       Zy)Ú_FeatureExtractorr%   Úreturnc                  ó   — y)áX  Generates features and length output from the given input tensor.

        Args:
            input (torch.Tensor): input tensor.

        Returns:
            (torch.Tensor, torch.Tensor):
            torch.Tensor:
                Features, with shape `(length, *)`.
            torch.Tensor:
                Length, with shape `(1,)`.
        N© r$   s     r   Ú__call__z_FeatureExtractor.__call__2   ó   � r   N)r(   r)   r*   r   r   ÚTensorr	   rB   rA   r   r   r=   r=   1   s8   „ Øð˜eŸl™lð ¨u°U·\±\À5Ç<Á<Ð5OÑ/Pò ó ñr   r=   c                   ó,   — e Zd Zedee   defd„«       Zy)Ú_TokenProcessorÚtokensr>   c                  ó   — y)zÊDecodes given list of tokens to text sequence.

        Args:
            tokens (List[int]): list of tokens to decode.

        Returns:
            str:
                Decoded text sequence.
        NrA   )r!   rG   Úkwargss      r   rB   z_TokenProcessor.__call__C   rC   r   N)r(   r)   r*   r   r   ÚintÚstrrB   rA   r   r   rF   rF   B   s&   „ Øð	˜t C™yð 	°sò 	ó ñ	r   rF   c                   óª   ‡ — e Zd ZdZdej
                  j                  ddfˆ fd„Zdej                  de	ej                  ej                  f   fd„Z
ˆ xZS )Ú_ModuleFeatureExtractorz›``torch.nn.Module``-based feature extraction pipeline.

    Args:
        pipeline (torch.nn.Module): module that implements feature extraction logic.
    Úpipeliner>   Nc                 ó0   •— t         ‰| �  «        || _        y r   )r   r   rN   )r!   rN   r"   s     €r   r   z _ModuleFeatureExtractor.__init__W   s   ø€ Ü‰ÑÔØ ˆ�r   r%   c                 ór   — | j                  |«      }t        j                  |j                  d   g«      }||fS )r@   r   )rN   r   r7   Úshape)r!   r%   ÚfeaturesÚlengths       r   r&   z_ModuleFeatureExtractor.forward[   s7   € ð —=‘= Ó'ˆÜ—‘˜xŸ~™~¨aÑ0Ð1Ó2ˆØ˜ÐÐr   )r(   r)   r*   Ú__doc__r   ÚnnÚModuler   rD   r	   r&   r+   r,   s   @r   rM   rM   P   sL   ø„ ñð! §¡§¡ð !°Tõ !ð ˜UŸ\™\ð  ¨e°E·L±LÀ%Ç,Á,Ð4NÑ.O÷  r   rM   c                   ó<   — e Zd ZdZdeddfd„Zd	dee   dedefd„Z	y)
Ú_SentencePieceTokenProcessorztSentencePiece-model-based token processor.

    Args:
        sp_model_path (str): path to SentencePiece model.
    Úsp_model_pathr>   Nc                 ó  — t        j                  d«      st        d«      ‚dd l}|j	                  |¬«      | _        | j
                  j                  «       | j
                  j                  «       | j
                  j                  «       h| _	        y )NÚsentencepiecez2SentencePiece is not available. Please install it.r   )Ú
model_file)
r
   Úis_module_availableÚRuntimeErrorr[   ÚSentencePieceProcessorÚsp_modelÚunk_idÚeos_idÚpad_idÚpost_process_remove_list)r!   rY   Úspms      r   r   z%_SentencePieceTokenProcessor.__init__t   sn   € Ü×/Ñ/°Ô@ÜÐSÓTÐTã#à×2Ñ2¸mÐ2ÓLˆŒà�M‰M× Ñ Ó"Ø�M‰M× Ñ Ó"Ø�M‰M× Ñ Ó"ð)
ˆÕ%r   rG   Úlstripc                 óê   — |dd D �cg c]  }|| j                   vsŒ|‘Œ }}dj                  | j                  j                  |«      «      j	                  dd«      }|r|j                  «       S |S c c}w )aX  Decodes given list of tokens to text sequence.

        Args:
            tokens (List[int]): list of tokens to decode.
            lstrip (bool, optional): if ``True``, returns text sequence with leading whitespace
                removed. (Default: ``True``).

        Returns:
            str:
                Decoded text sequence.
        é   NÚ u   â–�ú )rd   Újoinr`   Úid_to_pieceÚreplacerf   )r!   rG   rf   Útoken_indexÚfiltered_hypo_tokensÚoutput_strings         r   rB   z%_SentencePieceTokenProcessor.__call__�   s~   € ð ,2°!°"¨:ö 
Ø'¸ÈD×LiÑLiÒ9iŠKð 
Ðð  
ð Ÿ™ §¡× 9Ñ 9Ð:NÓ OÓP×XÑXÐYaÐcfÓgˆáØ ×'Ñ'Ó)Ð)à Ð ùò 
s
   ˆA0œA0)T)
r(   r)   r*   rT   rK   r   r   rJ   ÚboolrB   rA   r   r   rX   rX   m   s8   „ ñð
 cð 
¨dó 
ñ!˜t C™yð !°$ð !À#ô !r   rX   c                   ó€  — e Zd ZU dZ G d„ de«      Z G d„ de«      Zee	d<   e
g ef   e	d<   ee	d<   ee	d	<   ee	d
<   ee	d<   ee	d<   ee	d<   ee	d<   ee	d<   ee	d<   ee	d<   defd„Zedefd„«       Zedefd„«       Zedefd„«       Zedefd„«       Zedefd„«       Zedefd„«       Zdefd„Zdefd„Zdefd„Zdefd„Zy)Ú
RNNTBundleu«  Dataclass that bundles components for performing automatic speech recognition (ASR, speech-to-text)
    inference with an RNN-T model.

    More specifically, the class provides methods that produce the featurization pipeline,
    decoder wrapping the specified RNN-T model, and output token post-processor that together
    constitute a complete end-to-end ASR inference pipeline that produces a text sequence
    given a raw waveform.

    It can support non-streaming (full-context) inference as well as streaming inference.

    Users should not directly instantiate objects of this class; rather, users should use the
    instances (representing pre-trained models) that exist within the module,
    e.g. :data:`torchaudio.pipelines.EMFORMER_RNNT_BASE_LIBRISPEECH`.

    Example
        >>> import torchaudio
        >>> from torchaudio.pipelines import EMFORMER_RNNT_BASE_LIBRISPEECH
        >>> import torch
        >>>
        >>> # Non-streaming inference.
        >>> # Build feature extractor, decoder with RNN-T model, and token processor.
        >>> feature_extractor = EMFORMER_RNNT_BASE_LIBRISPEECH.get_feature_extractor()
        100%|â–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆ| 3.81k/3.81k [00:00<00:00, 4.22MB/s]
        >>> decoder = EMFORMER_RNNT_BASE_LIBRISPEECH.get_decoder()
        Downloading: "https://download.pytorch.org/torchaudio/models/emformer_rnnt_base_librispeech.pt"
        100%|â–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆ| 293M/293M [00:07<00:00, 42.1MB/s]
        >>> token_processor = EMFORMER_RNNT_BASE_LIBRISPEECH.get_token_processor()
        100%|â–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆâ–ˆ| 295k/295k [00:00<00:00, 25.4MB/s]
        >>>
        >>> # Instantiate LibriSpeech dataset; retrieve waveform for first sample.
        >>> dataset = torchaudio.datasets.LIBRISPEECH("/home/librispeech", url="test-clean")
        >>> waveform = next(iter(dataset))[0].squeeze()
        >>>
        >>> with torch.no_grad():
        >>>     # Produce mel-scale spectrogram features.
        >>>     features, length = feature_extractor(waveform)
        >>>
        >>>     # Generate top-10 hypotheses.
        >>>     hypotheses = decoder(features, length, 10)
        >>>
        >>> # For top hypothesis, convert predicted tokens to text.
        >>> text = token_processor(hypotheses[0][0])
        >>> print(text)
        he hoped there would be stew for dinner turnips and carrots and bruised potatoes and fat mutton pieces to [...]
        >>>
        >>>
        >>> # Streaming inference.
        >>> hop_length = EMFORMER_RNNT_BASE_LIBRISPEECH.hop_length
        >>> num_samples_segment = EMFORMER_RNNT_BASE_LIBRISPEECH.segment_length * hop_length
        >>> num_samples_segment_right_context = (
        >>>     num_samples_segment + EMFORMER_RNNT_BASE_LIBRISPEECH.right_context_length * hop_length
        >>> )
        >>>
        >>> # Build streaming inference feature extractor.
        >>> streaming_feature_extractor = EMFORMER_RNNT_BASE_LIBRISPEECH.get_streaming_feature_extractor()
        >>>
        >>> # Process same waveform as before, this time sequentially across overlapping segments
        >>> # to simulate streaming inference. Note the usage of ``streaming_feature_extractor`` and ``decoder.infer``.
        >>> state, hypothesis = None, None
        >>> for idx in range(0, len(waveform), num_samples_segment):
        >>>     segment = waveform[idx: idx + num_samples_segment_right_context]
        >>>     segment = torch.nn.functional.pad(segment, (0, num_samples_segment_right_context - len(segment)))
        >>>     with torch.no_grad():
        >>>         features, length = streaming_feature_extractor(segment)
        >>>         hypotheses, state = decoder.infer(features, length, 10, state=state, hypothesis=hypothesis)
        >>>     hypothesis = hypotheses[0]
        >>>     transcript = token_processor(hypothesis[0])
        >>>     if transcript:
        >>>         print(transcript, end=" ", flush=True)
        he hoped there would be stew for dinner turn ips and car rots and bru 'd oes and fat mut ton pieces to [...]
    c                   ó   — e Zd ZdZy)úRNNTBundle.FeatureExtractorz:Interface of the feature extraction part of RNN-T pipelineN©r(   r)   r*   rT   rA   r   r   ÚFeatureExtractorru   â   s   „ ÚHr   rw   c                   ó   — e Zd ZdZy)úRNNTBundle.TokenProcessorz7Interface of the token processor part of RNN-T pipelineNrv   rA   r   r   ÚTokenProcessorry   å   s   „ ÚEr   rz   Ú
_rnnt_pathÚ_rnnt_factory_funcÚ_global_stats_pathÚ_sp_model_pathÚ_right_paddingÚ_blankÚ_sample_rateÚ_n_fftÚ_n_melsÚ_hop_lengthÚ_segment_lengthÚ_right_context_lengthr>   c                 óä   — | j                  «       }t        j                  j                  | j                  «      }t        j                  |«      }|j                  |«       |j                  «        |S r   )	r|   Ú
torchaudioÚutilsÚdownload_assetr{   r   ÚloadÚload_state_dictÚeval)r!   ÚmodelÚpathÚ
state_dicts       r   Ú
_get_modelzRNNTBundle._get_modelõ   sT   € Ø×'Ñ'Ó)ˆÜ×Ñ×.Ñ.¨t¯©Ó?ˆÜ—Z‘Z Ó%ˆ
Ø×Ñ˜jÔ)Ø�
‰
ŒØˆr   c                 ó   — | j                   S )zSSample rate (in cycles per second) of input waveforms.

        :type: int
        )r�   ©r!   s    r   Úsample_ratezRNNTBundle.sample_rateý   s   € ð × Ñ Ð r   c                 ó   — | j                   S )z7Size of FFT window to use.

        :type: int
        )r‚   r“   s    r   Ún_fftzRNNTBundle.n_fft  s   € ð �{‰{Ðr   c                 ó   — | j                   S )z`Number of mel spectrogram features to extract from input waveforms.

        :type: int
        )rƒ   r“   s    r   Ún_melszRNNTBundle.n_mels  s   € ð �|‰|Ðr   c                 ó   — | j                   S )zdNumber of samples between successive frames in input expected by model.

        :type: int
        )r„   r“   s    r   Ú
hop_lengthzRNNTBundle.hop_length  s   € ð ×ÑÐr   c                 ó   — | j                   S )zTNumber of frames in segment in input expected by model.

        :type: int
        )r…   r“   s    r   Úsegment_lengthzRNNTBundle.segment_length  s   € ð ×#Ñ#Ð#r   c                 ó   — | j                   S )zcNumber of frames in right contextual block in input expected by model.

        :type: int
        )r†   r“   s    r   Úright_context_lengthzRNNTBundle.right_context_length%  s   € ð ×)Ñ)Ð)r   c                 óN   — | j                  «       }t        || j                  «      S )zOConstructs RNN-T decoder.

        Returns:
            RNNTBeamSearch
        )r‘   r   r€   )r!   rŽ   s     r   Úget_decoderzRNNTBundle.get_decoder-  s!   € ð —‘Ó!ˆÜ˜e T§[¡[Ó1Ð1r   c                 ó’  ‡ — t         j                  j                  ‰ j                  «      }t	        t
        j                  j                  t         j                  j                  ‰ j                  ‰ j                  ‰ j                  ‰ j                  ¬«      t        d„ «      t        d„ «      t        |«      t        ˆ fd„«      «      «      S )zzConstructs feature extractor for non-streaming (full-context) ASR.

        Returns:
            FeatureExtractor
        ©r”   r–   r˜   rš   c                 ó&   — | j                  dd«      S ©Nrh   r   ©Ú	transposer   s    r   ú<lambda>z2RNNTBundle.get_feature_extractor.<locals>.<lambda>B  ó   € ¨A¯K©K¸¸1Ó,=€ r   c                 ó&   — t        | t        z  «      S r   ©r   Ú_gainr   s    r   r§   z2RNNTBundle.get_feature_extractor.<locals>.<lambda>C  ó   € Ô,AÀ!ÄeÁ)Ó,L€ r   c                 ót   •— t         j                  j                  j                  | ddd‰j                  f«      S )Nr   )r   rU   r    Úpadr   )r   r!   s    €r   r§   z2RNNTBundle.get_feature_extractor.<locals>.<lambda>E  s.   ø€ ¬E¯H©H×,?Ñ,?×,CÑ,CÀAÈÈ1ÈaÐQU×QdÑQdÐGeÓ,f€ r   ©rˆ   r‰   rŠ   r}   rM   r   rU   Ú
SequentialÚ
transformsÚMelSpectrogramr”   r–   r˜   rš   r   r.   ©r!   Ú
local_paths   ` r   Úget_feature_extractorz RNNTBundle.get_feature_extractor6  sœ   ø€ ô  ×%Ñ%×4Ñ4°T×5LÑ5LÓMˆ
Ü&Ü�H‰H×ÑÜ×%Ñ%×4Ñ4Ø $× 0Ñ 0¸¿
¹
È4Ï;É;Ðcg×crÑcrð 5ó ô "Ñ"=Ó>Ü!Ñ"LÓMÜ)¨*Ó5Ü!Ó"fÓgóó

ð 
	
r   c           
      óv  — t         j                  j                  | j                  «      }t	        t
        j                  j                  t         j                  j                  | j                  | j                  | j                  | j                  ¬«      t        d„ «      t        d„ «      t        |«      «      «      S )zvConstructs feature extractor for streaming (simultaneous) ASR.

        Returns:
            FeatureExtractor
        r¢   c                 ó&   — | j                  dd«      S r¤   r¥   r   s    r   r§   z<RNNTBundle.get_streaming_feature_extractor.<locals>.<lambda>U  r¨   r   c                 ó&   — t        | t        z  «      S r   rª   r   s    r   r§   z<RNNTBundle.get_streaming_feature_extractor.<locals>.<lambda>V  r¬   r   r¯   r³   s     r   Úget_streaming_feature_extractorz*RNNTBundle.get_streaming_feature_extractorI  s’   € ô  ×%Ñ%×4Ñ4°T×5LÑ5LÓMˆ
Ü&Ü�H‰H×ÑÜ×%Ñ%×4Ñ4Ø $× 0Ñ 0¸¿
¹
È4Ï;É;Ðcg×crÑcrð 5ó ô "Ñ"=Ó>Ü!Ñ"LÓMÜ)¨*Ó5óó	
ð 		
r   c                 ój   — t         j                  j                  | j                  «      }t	        |«      S )zQConstructs token processor.

        Returns:
            TokenProcessor
        )rˆ   r‰   rŠ   r~   rX   r³   s     r   Úget_token_processorzRNNTBundle.get_token_processor[  s+   € ô  ×%Ñ%×4Ñ4°T×5HÑ5HÓIˆ
Ü+¨JÓ7Ð7r   N)r(   r)   r*   rT   r=   rw   rF   rz   rK   Ú__annotations__r   r   rJ   r‘   Úpropertyr”   r–   r˜   rš   rœ   rž   r   r    rµ   r¹   r»   rA   r   r   rs   rs   ˜   sU  … ñFôPIÐ,ô IôF˜ô Fð ƒOØ   T Ñ*Ó*ØÓØÓØÓØƒKØÓØƒKØƒLØÓØÓØÓð˜Dó ð ð!˜Sò !ó ð!ð ð�sò ó ðð ð˜ò ó ðð ð ˜Cò  ó ð ð ð$ ò $ó ð$ð ð* cò *ó ð*ð2˜^ó 2ð
Ð'7ó 
ð&
Ð1Aó 
ð$8 ^ô 8r   rs   z(models/emformer_rnnt_base_librispeech.pti  )Únum_symbolsz2pipeline-assets/global_stats_rnnt_librispeech.jsonz.pipeline-assets/spm_bpe_4096_librispeech.modelé   i   i€>  i�  éP   é    é   )r{   r|   r}   r~   r   r€   r�   r‚   rƒ   r„   r…   r†   aì  ASR pipeline based on Emformer-RNNT,
pretrained on *LibriSpeech* dataset :cite:`7178964`,
capable of performing both streaming and non-streaming inference.

The underlying model is constructed by :py:func:`torchaudio.models.emformer_rnnt_base`
and utilizes weights trained on LibriSpeech using training script ``train.py``
`here <https://github.com/pytorch/audio/tree/main/examples/asr/emformer_rnnt>`__ with default arguments.

Please refer to :py:class:`RNNTBundle` for usage instructions.
))r3   r   Úabcr   r   Údataclassesr   Ú	functoolsr   Útypingr   r   r	   r   rˆ   Útorchaudio._internalr
   Útorchaudio.modelsr   r   r   Ú__all__Úlog10ÚiinfoÚint16ÚmaxÚ_decibelÚpowr«   r   rU   rV   r   r.   r=   rF   rM   rX   rs   ÚEMFORMER_RNNT_BASE_LIBRISPEECHrT   rA   r   r   ú<module>rÑ      s4  ðÛ Û ß #Ý !Ý ß (Ñ (ã Û Ý -ß FÑ Fð €à�J�D—J‘J˜{˜uŸ{™{¨5¯;©;Ó7×;Ñ;Ó<Ñ<€ÙˆB��x‘Ó €òô&˜Ÿ™Ÿ™ô &ô4 §¡§¡ô 4ô˜ô ô"�cô ô ˜eŸh™hŸo™oÐ/@ô  ô:(! ?ô (!ðV ÷I8ð I8ó ðI8ñX ",Ø9ÙÐ1¸tÔDØKØCØØØØØØØØô"Ð ð	*Ð Õ &r   