Ë
    óÿæiÕŠ  ã            ,       ó@  — d dl mZmZ d dlmZmZmZ d dlZd dlm	Z	 g 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j                  j                  e«      Z G d„ dej                  j                  «      Z G d„ dej                  j                  «      Z G d„ dej                  j                  «      Zdedededededededededededed ed!ed"ed#ed$ed%ed&ed'ed(ed)ef,d*„Zded)efd+„Zy),é    )ÚABCÚabstractmethod)ÚListÚOptionalÚTupleN)ÚEmformer)ÚRNNTÚemformer_rnnt_baseÚemformer_rnnt_modelc                   óš   ‡ — e Zd ZdZdeddfˆ fd„Zdej                  dej                  deej                  ej                  f   fd„Z	ˆ xZ
S )	Ú_TimeReductionzÂCoalesces frames along time dimension into a
    fewer number of frames with higher feature dimensionality.

    Args:
        stride (int): number of frames to merge for each output frame.
    ÚstrideÚreturnNc                 ó0   •— t         ‰| �  «        || _        y ©N)ÚsuperÚ__init__r   )Úselfr   Ú	__class__s     €úk/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchaudio/models/rnnt.pyr   z_TimeReduction.__init__   s   ø€ Ü‰ÑÔØˆ�ó    ÚinputÚlengthsc                 ó"  — |j                   \  }}}||| j                  z  z
  }|dd…d|…dd…f   }|j                  | j                  d¬«      }|| j                  z  }|j                  |||| j                  z  «      }|j	                  «       }||fS )a  Forward pass.

        B: batch size;
        T: maximum input sequence length in batch;
        D: feature dimension of each input sequence frame.

        Args:
            input (torch.Tensor): input sequences, with shape `(B, T, D)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``input``.

        Returns:
            (torch.Tensor, torch.Tensor):
                torch.Tensor
                    output sequences, with shape
                    `(B, T  // stride, D * stride)`
                torch.Tensor
                    output lengths, with shape `(B,)` and i-th element representing
                    number of valid frames for i-th batch element in output sequences.
        NÚtrunc)Úrounding_mode)Úshaper   ÚdivÚreshapeÚ
contiguous)	r   r   r   ÚBÚTÚDÚ
num_framesÚT_maxÚoutputs	            r   Úforwardz_TimeReduction.forward   s�   € ð* —+‘+‰ˆˆ1ˆaØ˜!˜dŸk™k™/Ñ*ˆ
Ø’a˜˜*˜¢aÐ'Ñ(ˆØ—+‘+˜dŸk™k¸�+ÓAˆØ˜dŸk™kÑ)ˆà—‘˜q %¨¨T¯[©[©Ó9ˆØ×"Ñ"Ó$ˆØ�wˆÐr   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úintr   ÚtorchÚTensorr   r'   Ú__classcell__©r   s   @r   r   r      sR   ø„ ñð˜sð  tõ ð˜UŸ\™\ð °E·L±Lð ÀUÈ5Ï<É<ÐY^×YeÑYeÐKeÑEf÷ r   r   c                   ó¾   ‡ — e Zd ZdZ	 	 ddededededdf
ˆ fd„Zd	ej                  d
e
eej                        deej                  eej                     f   fd„Zˆ xZS )Ú_CustomLSTMa­  Custom long-short-term memory (LSTM) block that applies layer normalization
    to internal nodes.

    Args:
        input_dim (int): input dimension.
        hidden_dim (int): hidden dimension.
        layer_norm (bool, optional): if ``True``, enables layer normalization. (Default: ``False``)
        layer_norm_epsilon (float, optional):  value of epsilon to use in
            layer normalization layers (Default: 1e-5)
    Ú	input_dimÚ
hidden_dimÚ
layer_normÚlayer_norm_epsilonr   Nc                 ó  •— t         ‰| �  «        t        j                  j	                  |d|z  | ¬«      | _        t        j                  j	                  |d|z  d¬«      | _        |rWt        j                  j                  ||¬«      | _        t        j                  j                  d|z  |¬«      | _	        || _        y t        j                  j                  «       | _        t        j                  j                  «       | _	        || _        y )Né   ©ÚbiasF)Úeps)r   r   r-   ÚnnÚLinearÚx2gÚp2gÚ	LayerNormÚc_normÚg_normÚIdentityr4   )r   r3   r4   r5   r6   r   s        €r   r   z_CustomLSTM.__init__C   sÇ   ø€ ô 	‰ÑÔÜ—8‘8—?‘? 9¨a°*©nÈ
ÀN�?ÓTˆŒÜ—8‘8—?‘? :¨q°:©~ÀE�?ÓJˆŒÙÜŸ(™(×,Ñ,¨ZÐ=OÐ,ÓPˆDŒKÜŸ(™(×,Ñ,¨Q°©^ÐASÐ,ÓTˆDŒKð
 %ˆ�ô  Ÿ(™(×+Ñ+Ó-ˆDŒKÜŸ(™(×+Ñ+Ó-ˆDŒKà$ˆ�r   r   Ústatec                 ó  — |€€|j                  d«      }t        j                  || j                  |j                  |j
                  ¬«      }t        j                  || j                  |j                  |j
                  ¬«      }n|\  }}| j                  |«      }g }|j                  d«      D ]¾  }|| j                  |«      z   }| j                  |«      }|j                  dd«      \  }	}
}}|	j                  «       }	|
j                  «       }
|j                  «       }|j                  «       }|
|z  |	|z  z   }| j                  |«      }||j                  «       z  }|j                  |«       ŒÀ t        j                  |d¬«      }||g}||fS )aÒ  Forward pass.

        B: batch size;
        T: maximum sequence length in batch;
        D: feature dimension of each input sequence element.

        Args:
            input (torch.Tensor): with shape `(T, B, D)`.
            state (List[torch.Tensor] or None): list of tensors
                representing internal state generated in preceding invocation
                of ``forward``.

        Returns:
            (torch.Tensor, List[torch.Tensor]):
                torch.Tensor
                    output, with shape `(T, B, hidden_dim)`.
                List[torch.Tensor]
                    list of tensors representing internal state generated
                    in current invocation of ``forward``.
        é   )ÚdeviceÚdtyper   r8   )Údim)Úsizer-   Úzerosr4   rG   rH   r>   Úunbindr?   rB   ÚchunkÚsigmoidÚtanhrA   ÚappendÚstack)r   r   rD   r!   ÚhÚcÚgated_inputÚoutputsÚgatesÚ
input_gateÚforget_gateÚ	cell_gateÚoutput_gater&   s                 r   r'   z_CustomLSTM.forwardV   sV  € ð. ˆ=Ø—
‘
˜1“ˆAÜ—‘˜A˜tŸ™°u·|±|È5Ï;É;ÔWˆAÜ—‘˜A˜tŸ™°u·|±|È5Ï;É;ÔW‰Aà‰DˆAˆqà—h‘h˜u“oˆØˆØ ×'Ñ'¨Ö*ˆEØ˜DŸH™H Q›KÑ'ˆEØ—K‘K Ó&ˆEØ>C¿k¹kÈ!ÈQÓ>OÑ;ˆJ˜ Y°Ø#×+Ñ+Ó-ˆJØ%×-Ñ-Ó/ˆKØ!Ÿ™Ó(ˆIØ%×-Ñ-Ó/ˆKØ˜a‘ *¨yÑ"8Ñ8ˆAØ—‘˜A“ˆAØ˜aŸf™f›hÑ&ˆAØ�N‰N˜1Õð +ô —‘˜W¨!Ô,ˆØ�A�ˆà�uˆ}Ðr   )Fçñhãˆµøä>©r(   r)   r*   r+   r,   ÚboolÚfloatr   r-   r.   r   r   r   r'   r/   r0   s   @r   r2   r2   7   s‹   ø„ ñ	ð !Ø$(ñ%àð%ð ð%ð ð	%ð
 "ð%ð 
õ%ð&0Ø—\‘\ð0Ø*2°4¸¿¹Ñ3EÑ*Fð0à	ˆu�|‰|˜T %§,¡,Ñ/Ð/Ñ	0÷0r   r2   c                   óH  — e Zd Zedej
                  dej
                  deej
                  ej
                  f   fd„«       Zedej
                  dej
                  dee	e	ej
                           deej
                  ej
                  e	e	ej
                        f   fd„«       Z
y)Ú_Transcriberr   r   r   c                  ó   — y r   © )r   r   r   s      r   r'   z_Transcriber.forwardŠ   s   € àr   Ústatesc                  ó   — y r   rb   )r   r   r   rc   s       r   Úinferz_Transcriber.inferŽ   s   € ð 	r   N)r(   r)   r*   r   r-   r.   r   r'   r   r   re   rb   r   r   r`   r`   ‰   s½   „ Øð˜UŸ\™\ð °E·L±Lð ÀUÈ5Ï<É<ÐY^×YeÑYeÐKeÑEfò ó ðð ðà�|‰|ðð —‘ðð ˜˜d 5§<¡<Ñ0Ñ1Ñ2ð	ð
 
ˆu�|‰|˜UŸ\™\¨4°°U·\±\Ñ0BÑ+CÐCÑ	Dòó ñr   r`   c            !       óÌ  ‡ — e Zd ZdZddddddœded	ed
edededededededededededededdf ˆ fd„Zde	j                  de	j                  dee	j                  e	j                  f   fd„Ze	j                  j                  de	j                  de	j                  deeee	j                           dee	j                  e	j                  eee	j                        f   fd„«       Zˆ xZS )Ú_EmformerEncodera½  Emformer-based recurrent neural network transducer (RNN-T) encoder (transcription network).

    Args:
        input_dim (int): feature dimension of each input sequence element.
        output_dim (int): feature dimension of each output sequence element.
        segment_length (int): length of input segment expressed as number of frames.
        right_context_length (int): length of right context expressed as number of frames.
        time_reduction_input_dim (int): dimension to scale each element in input sequences to
            prior to applying time reduction block.
        time_reduction_stride (int): factor by which to reduce length of input sequence.
        transformer_num_heads (int): number of attention heads in each Emformer layer.
        transformer_ffn_dim (int): hidden layer dimension of each Emformer layer's feedforward network.
        transformer_num_layers (int): number of Emformer layers to instantiate.
        transformer_left_context_length (int): length of left context.
        transformer_dropout (float, optional): transformer dropout probability. (Default: 0.0)
        transformer_activation (str, optional): activation function to use in each Emformer layer's
            feedforward network. Must be one of ("relu", "gelu", "silu"). (Default: "relu")
        transformer_max_memory_size (int, optional): maximum number of memory elements to use. (Default: 0)
        transformer_weight_init_scale_strategy (str, optional): per-layer weight initialization scaling
            strategy. Must be one of ("depthwise", "constant", ``None``). (Default: "depthwise")
        transformer_tanh_on_mem (bool, optional): if ``True``, applies tanh to memory elements. (Default: ``False``)
    ç        Úrelur   Ú	depthwiseF)Útransformer_dropoutÚtransformer_activationÚtransformer_max_memory_sizeÚ&transformer_weight_init_scale_strategyÚtransformer_tanh_on_memr3   Ú
output_dimÚsegment_lengthÚright_context_lengthÚtime_reduction_input_dimÚtime_reduction_strideÚtransformer_num_headsÚtransformer_ffn_dimÚtransformer_num_layersÚtransformer_left_context_lengthrk   rl   rm   rn   ro   r   Nc                óp  •— t         ‰| �  «        t        j                  j	                  ||d¬«      | _        t        |«      | _        ||z  }t        ||||	||z  |||
||z  |||¬«      | _	        t        j                  j	                  ||«      | _
        t        j                  j                  |«      | _        y )NFr9   )ÚdropoutÚ
activationÚleft_context_lengthrr   Úmax_memory_sizeÚweight_init_scale_strategyÚtanh_on_mem)r   r   r-   r<   r=   Úinput_linearr   Útime_reductionr   ÚtransformerÚoutput_linearr@   r5   )r   r3   rp   rq   rr   rs   rt   ru   rv   rw   rx   rk   rl   rm   rn   ro   Útransformer_input_dimr   s                    €r   r   z_EmformerEncoder.__init__°   s¼   ø€ ô& 	‰ÑÔÜ!ŸH™HŸO™OØØ$Øð ,ó 
ˆÔô
 -Ð-BÓCˆÔØ 8Ð;PÑ PÐÜ#Ø!Ø!ØØ"ØÐ3Ñ3Ø'Ø-Ø ?Ø!5Ð9NÑ!NØ7Ø'MØ/ô
ˆÔô #ŸX™XŸ_™_Ð-BÀJÓOˆÔÜŸ(™(×,Ñ,¨ZÓ8ˆ�r   r   r   c                 óÄ   — | j                  |«      }| j                  ||«      \  }}| j                  ||«      \  }}| j                  |«      }| j	                  |«      }	|	|fS )a¥  Forward pass for training.

        B: batch size;
        T: maximum input sequence length in batch;
        D: feature dimension of each input sequence frame (input_dim).

        Args:
            input (torch.Tensor): input frame sequences right-padded with right context, with
                shape `(B, T + right context length, D)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``input``.

        Returns:
            (torch.Tensor, torch.Tensor):
                torch.Tensor
                    output frame sequences, with
                    shape `(B, T // time_reduction_stride, output_dim)`.
                torch.Tensor
                    output input lengths, with shape `(B,)` and i-th element representing
                    number of valid elements for i-th batch element in output frame sequences.
        )r€   r�   r‚   rƒ   r5   )
r   r   r   Úinput_linear_outÚtime_reduction_outÚtime_reduction_lengthsÚtransformer_outÚtransformer_lengthsÚoutput_linear_outÚlayer_norm_outs
             r   r'   z_EmformerEncoder.forwardÜ   sv   € ð,  ×,Ñ,¨UÓ3ÐØ59×5HÑ5HÐIYÐ[bÓ5cÑ2ÐÐ2Ø/3×/?Ñ/?Ð@RÐTjÓ/kÑ,ˆÐ,Ø ×.Ñ.¨Ó?ÐØŸ™Ð):Ó;ˆØÐ2Ð2Ð2r   rc   c                 óÞ   — | j                  |«      }| j                  ||«      \  }}| j                  j                  |||«      \  }}}	| j	                  |«      }
| j                  |
«      }|||	fS )aR  Forward pass for inference.

        B: batch size;
        T: maximum input sequence segment length in batch;
        D: feature dimension of each input sequence frame (input_dim).

        Args:
            input (torch.Tensor): input frame sequence segments right-padded with right context, with
                shape `(B, T + right context length, D)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``input``.
            state (List[List[torch.Tensor]] or None): list of lists of tensors
                representing internal state generated in preceding invocation
                of ``infer``.

        Returns:
            (torch.Tensor, torch.Tensor, List[List[torch.Tensor]]):
                torch.Tensor
                    output frame sequences, with
                    shape `(B, T // time_reduction_stride, output_dim)`.
                torch.Tensor
                    output input lengths, with shape `(B,)` and i-th element representing
                    number of valid elements for i-th batch element in output.
                List[List[torch.Tensor]]
                    output states; list of lists of tensors
                    representing internal state generated in current invocation
                    of ``infer``.
        )r€   r�   r‚   re   rƒ   r5   )r   r   r   rc   r†   r‡   rˆ   r‰   rŠ   Útransformer_statesr‹   rŒ   s               r   re   z_EmformerEncoder.inferù   sŠ   € ðF  ×,Ñ,¨UÓ3ÐØ59×5HÑ5HÐIYÐ[bÓ5cÑ2ÐÐ2ð
 ×Ñ×"Ñ"Ð#5Ð7MÈvÓVñ		
ØØØà ×.Ñ.¨Ó?ÐØŸ™Ð):Ó;ˆØÐ2Ð4FÐFÐFr   )r(   r)   r*   r+   r,   r^   Ústrr]   r   r-   r.   r   r'   ÚjitÚexportr   r   re   r/   r0   s   @r   rg   rg   ˜   s~  ø„ ñðH &)Ø&,Ø+,Ø6AØ(-ò#*9ð ð*9ð ð	*9ð
 ð*9ð "ð*9ð #&ð*9ð  #ð*9ð  #ð*9ð !ð*9ð !$ð*9ð *-ð*9ð #ð*9ð !$ð*9ð &)ð*9ð  14ð!*9ð" "&ð#*9ð$ 
õ%*9ðX3˜UŸ\™\ð 3°E·L±Lð 3ÀUÈ5Ï<É<ÐY^×YeÑYeÐKeÑEfó 3ð: ‡Y�Y×Ñð+Gà�|‰|ð+Gð —‘ð+Gð ˜˜d 5§<¡<Ñ0Ñ1Ñ2ð	+Gð
 
ˆu�|‰|˜UŸ\™\¨4°°U·\±\Ñ0BÑ+CÐCÑ	Dò+Gó ô+Gr   rg   c                   ó  ‡ — e Zd ZdZ	 	 	 ddededededededed	ed
dfˆ fd„Z	 ddej                  dej                  de
eeej                           d
eej                  ej                  eeej                        f   fd„Zˆ xZS )Ú
_Predictora  Recurrent neural network transducer (RNN-T) prediction network.

    Args:
        num_symbols (int): size of target token lexicon.
        output_dim (int): feature dimension of each output sequence element.
        symbol_embedding_dim (int): dimension of each target token embedding.
        num_lstm_layers (int): number of LSTM layers to instantiate.
        lstm_hidden_dim (int): output dimension of each LSTM layer.
        lstm_layer_norm (bool, optional): if ``True``, enables layer normalization
            for LSTM layers. (Default: ``False``)
        lstm_layer_norm_epsilon (float, optional): value of epsilon to use in
            LSTM layer normalization layers. (Default: 1e-5)
        lstm_dropout (float, optional): LSTM dropout probability. (Default: 0.0)

    Únum_symbolsrp   Úsymbol_embedding_dimÚnum_lstm_layersÚlstm_hidden_dimÚlstm_layer_normÚlstm_layer_norm_epsilonÚlstm_dropoutr   Nc	                 óF  •— t         ‰
| �  «        t        j                  j	                  ||«      | _        t        j                  j                  |«      | _        t        j                  j                  t        |«      D �	cg c]  }	t        |	dk(  r|n||||¬«      ‘Œ c}	«      | _        t        j                  j                  |¬«      | _        t        j                  j                  ||«      | _        t        j                  j                  |«      | _        || _        y c c}	w )Nr   )r5   r6   )Úp)r   r   r-   r<   Ú	EmbeddingÚ	embeddingr@   Úinput_layer_normÚ
ModuleListÚranger2   Úlstm_layersÚDropoutrz   r=   ÚlinearÚoutput_layer_normrš   )r   r”   rp   r•   r–   r—   r˜   r™   rš   Úidxr   s             €r   r   z_Predictor.__init__9  sí   ø€ ô 	‰ÑÔÜŸ™×+Ñ+¨KÐ9MÓNˆŒÜ %§¡× 2Ñ 2Ð3GÓ HˆÔÜ Ÿ8™8×.Ñ.ô ! Ô1óñ 2�Cô Ø,/°1ªHÑ(¸/Ø#Ø.Ø'>ö	ð 2ñó

ˆÔô —x‘x×'Ñ'¨,Ð'Ó7ˆŒÜ—h‘h—o‘o o°zÓBˆŒÜ!&§¡×!3Ñ!3°JÓ!?ˆÔà(ˆÕùòs   Á?Dr   r   rD   c                 ó†  — |j                  dd«      }| j                  |«      }| j                  |«      }|}g }t        | j                  «      D ]:  \  }	}
 |
||€dn||	   «      \  }}| j                  |«      }|j                  |«       Œ< | j                  |«      }| j                  |«      }|j                  ddd«      ||fS )a#  Forward pass.

        B: batch size;
        U: maximum sequence length in batch;
        D: feature dimension of each input sequence element.

        Args:
            input (torch.Tensor): target sequences, with shape `(B, U)` and each element
                mapping to a target symbol, i.e. in range `[0, num_symbols)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``input``.
            state (List[List[torch.Tensor]] or None, optional): list of lists of tensors
                representing internal state generated in preceding invocation
                of ``forward``. (Default: ``None``)

        Returns:
            (torch.Tensor, torch.Tensor, List[List[torch.Tensor]]):
                torch.Tensor
                    output encoding sequences, with shape `(B, U, output_dim)`
                torch.Tensor
                    output lengths, with shape `(B,)` and i-th element representing
                    number of valid elements for i-th batch element in output encoding sequences.
                List[List[torch.Tensor]]
                    output states; list of lists of tensors
                    representing internal state generated in current invocation of ``forward``.
        rF   r   Né   )	Úpermuterž   rŸ   Ú	enumerater¢   rz   rP   r¤   r¥   )r   r   r   rD   Úinput_tbÚembedding_outÚinput_layer_norm_outÚlstm_outÚ	state_outÚ	layer_idxÚlstmÚlstm_state_outÚ
linear_outÚoutput_layer_norm_outs                 r   r'   z_Predictor.forwardX  sÎ   € ð@ —=‘=  AÓ&ˆØŸ™ xÓ0ˆØ#×4Ñ4°]ÓCÐà'ˆØ.0ˆ	Ü(¨×)9Ñ)9Ö:‰OˆI�tÙ'+¨H¸e¸m±dÐQVÐW`ÑQaÓ'bÑ$ˆH�nØ—|‘| HÓ-ˆHØ×Ñ˜^Õ,ð  ;ð
 —[‘[ Ó*ˆ
Ø $× 6Ñ 6°zÓ BÐØ$×,Ñ,¨Q°°1Ó5°wÀ	ÐIÐIr   )Fr[   rh   r   r\   r0   s   @r   r“   r“   (  sã   ø„ ñð. !&Ø)-Ø!ñ)àð)ð ð)ð "ð	)ð
 ð)ð ð)ð ð)ð "'ð)ð ð)ð 
õ)ðF 59ñ	-Jà�|‰|ð-Jð —‘ð-Jð ˜˜T %§,¡,Ñ/Ñ0Ñ1ð	-Jð
 
ˆu�|‰|˜UŸ\™\¨4°°U·\±\Ñ0BÑ+CÐCÑ	D÷-Jr   r“   c                   óê   ‡ — e Zd ZdZddedededdfˆ fd„Zdej                  d	ej                  d
ej                  dej                  de	ej                  ej                  ej                  f   f
d„Z
ˆ xZS )Ú_Joinera@  Recurrent neural network transducer (RNN-T) joint network.

    Args:
        input_dim (int): source and target input dimension.
        output_dim (int): output dimension.
        activation (str, optional): activation function to use in the joiner.
            Must be one of ("relu", "tanh"). (Default: "relu")

    r3   rp   r{   r   Nc                 ó.  •— t         ‰| �  «        t        j                  j	                  ||d¬«      | _        |dk(  r$t        j                  j                  «       | _        y |dk(  r$t        j                  j                  «       | _        y t        d|› �«      ‚)NTr9   ri   rO   zUnsupported activation )
r   r   r-   r<   r=   r¤   ÚReLUr{   ÚTanhÚ
ValueError)r   r3   rp   r{   r   s       €r   r   z_Joiner.__init__“  sm   ø€ Ü‰ÑÔÜ—h‘h—o‘o i°À$�oÓGˆŒØ˜ÒÜ#Ÿh™hŸm™m›oˆD�OØ˜6Ò!Ü#Ÿh™hŸm™m›oˆD�OäÐ6°z°lÐCÓDÐDr   Úsource_encodingsÚsource_lengthsÚtarget_encodingsÚtarget_lengthsc                 óÎ   — |j                  d«      j                  «       |j                  d«      j                  «       z   }| j                  |«      }| j                  |«      }|||fS )a›  Forward pass for training.

        B: batch size;
        T: maximum source sequence length in batch;
        U: maximum target sequence length in batch;
        D: dimension of each source and target sequence encoding.

        Args:
            source_encodings (torch.Tensor): source encoding sequences, with
                shape `(B, T, D)`.
            source_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                valid sequence length of i-th batch element in ``source_encodings``.
            target_encodings (torch.Tensor): target encoding sequences, with shape `(B, U, D)`.
            target_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                valid sequence length of i-th batch element in ``target_encodings``.

        Returns:
            (torch.Tensor, torch.Tensor, torch.Tensor):
                torch.Tensor
                    joint network output, with shape `(B, T, U, output_dim)`.
                torch.Tensor
                    output source lengths, with shape `(B,)` and i-th element representing
                    number of valid elements along dim 1 for i-th batch element in joint network output.
                torch.Tensor
                    output target lengths, with shape `(B,)` and i-th element representing
                    number of valid elements along dim 2 for i-th batch element in joint network output.
        r¨   rF   )Ú	unsqueezer    r{   r¤   )r   r»   r¼   r½   r¾   Újoint_encodingsÚactivation_outr&   s           r   r'   z_Joiner.forward�  sb   € ðD +×4Ñ4°QÓ7×BÑBÓDÐGW×GaÑGaÐbcÓGd×GoÑGoÓGqÑqˆØŸ™¨Ó9ˆØ—‘˜^Ó,ˆØ�~ ~Ð5Ð5r   )ri   )r(   r)   r*   r+   r,   r�   r   r-   r.   r   r'   r/   r0   s   @r   r¶   r¶   ˆ  sŒ   ø„ ññE #ð E°3ð EÀCð EÐUYõ Eð%6àŸ,™,ð%6ð Ÿ™ð%6ð  Ÿ,™,ð	%6ð
 Ÿ™ð%6ð 
ˆu�|‰|˜UŸ\™\¨5¯<©<Ð7Ñ	8÷%6r   r¶   c                   ó–  ‡ — e Zd ZdZdedededdfˆ fd„Z	 ddej                  d	ej                  d
ej                  dej                  de
eeej                           deej                  ej                  ej                  eeej                        f   fd„Zej                  j                  dej                  d	ej                  de
eeej                           deej                  ej                  eeej                        f   fd„«       Zej                  j                  dej                  d	ej                  deej                  ej                  f   fd„«       Zej                  j                  d
ej                  dej                  de
eeej                           deej                  ej                  eeej                        f   fd„«       Zej                  j                  dej                  d	ej                  dej                  dej                  deej                  ej                  ej                  f   f
d„«       Zˆ xZS )r	   a¿  torchaudio.models.RNNT()

    Recurrent neural network transducer (RNN-T) model.

    Note:
        To build the model, please use one of the factory functions.

    See Also:
        :class:`torchaudio.pipelines.RNNTBundle`: ASR pipeline with pre-trained models.

    Args:
        transcriber (torch.nn.Module): transcription network.
        predictor (torch.nn.Module): prediction network.
        joiner (torch.nn.Module): joint network.
    ÚtranscriberÚ	predictorÚjoinerr   Nc                 óL   •— t         ‰| �  «        || _        || _        || _        y r   )r   r   rÄ   rÅ   rÆ   )r   rÄ   rÅ   rÆ   r   s       €r   r   zRNNT.__init__Ö  s$   ø€ Ü‰ÑÔØ&ˆÔØ"ˆŒØˆ�r   Úsourcesr¼   Útargetsr¾   Úpredictor_statec                 óœ   — | j                  ||¬«      \  }}| j                  |||¬«      \  }}}| j                  ||||¬«      \  }}}||||fS )a  Forward pass for training.

        B: batch size;
        T: maximum source sequence length in batch;
        U: maximum target sequence length in batch;
        D: feature dimension of each source sequence element.

        Args:
            sources (torch.Tensor): source frame sequences right-padded with right context, with
                shape `(B, T, D)`.
            source_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``sources``.
            targets (torch.Tensor): target sequences, with shape `(B, U)` and each element
                mapping to a target symbol.
            target_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``targets``.
            predictor_state (List[List[torch.Tensor]] or None, optional): list of lists of tensors
                representing prediction network internal state generated in preceding invocation
                of ``forward``. (Default: ``None``)

        Returns:
            (torch.Tensor, torch.Tensor, torch.Tensor, List[List[torch.Tensor]]):
                torch.Tensor
                    joint network output, with shape
                    `(B, max output source length, max output target length, output_dim (number of target symbols))`.
                torch.Tensor
                    output source lengths, with shape `(B,)` and i-th element representing
                    number of valid elements along dim 1 for i-th batch element in joint network output.
                torch.Tensor
                    output target lengths, with shape `(B,)` and i-th element representing
                    number of valid elements along dim 2 for i-th batch element in joint network output.
                List[List[torch.Tensor]]
                    output states; list of lists of tensors
                    representing prediction network internal state generated in current invocation
                    of ``forward``.
        )r   r   ©r   r   rD   ©r»   r¼   r½   r¾   )rÄ   rÅ   rÆ   )	r   rÈ   r¼   rÉ   r¾   rÊ   r»   r½   r&   s	            r   r'   zRNNT.forwardÜ  sŽ   € ðX ,0×+;Ñ+;ØØ"ð ,<ó ,
Ñ(Ð˜.ð =A¿N¹NØØ"Ø!ð =Kó =
Ñ9Ð˜.¨/ð
 26·±Ø-Ø)Ø-Ø)ð	 2=ó 2
Ñ.ˆ� ð ØØØð	
ð 	
r   rD   c                 ó<   — | j                   j                  |||«      S )a¸  Applies transcription network to sources in streaming mode.

        B: batch size;
        T: maximum source sequence segment length in batch;
        D: feature dimension of each source sequence frame.

        Args:
            sources (torch.Tensor): source frame sequence segments right-padded with right context, with
                shape `(B, T + right context length, D)`.
            source_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``sources``.
            state (List[List[torch.Tensor]] or None): list of lists of tensors
                representing transcription network internal state generated in preceding invocation
                of ``transcribe_streaming``.

        Returns:
            (torch.Tensor, torch.Tensor, List[List[torch.Tensor]]):
                torch.Tensor
                    output frame sequences, with
                    shape `(B, T // time_reduction_stride, output_dim)`.
                torch.Tensor
                    output lengths, with shape `(B,)` and i-th element representing
                    number of valid elements for i-th batch element in output.
                List[List[torch.Tensor]]
                    output states; list of lists of tensors
                    representing transcription network internal state generated in current invocation
                    of ``transcribe_streaming``.
        )rÄ   re   )r   rÈ   r¼   rD   s       r   Útranscribe_streamingzRNNT.transcribe_streaming  s    € ðF ×Ñ×%Ñ% g¨~¸uÓEÐEr   c                 ó&   — | j                  ||«      S )aÆ  Applies transcription network to sources in non-streaming mode.

        B: batch size;
        T: maximum source sequence length in batch;
        D: feature dimension of each source sequence frame.

        Args:
            sources (torch.Tensor): source frame sequences right-padded with right context, with
                shape `(B, T + right context length, D)`.
            source_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``sources``.

        Returns:
            (torch.Tensor, torch.Tensor):
                torch.Tensor
                    output frame sequences, with
                    shape `(B, T // time_reduction_stride, output_dim)`.
                torch.Tensor
                    output lengths, with shape `(B,)` and i-th element representing
                    number of valid elements for i-th batch element in output frame sequences.
        )rÄ   )r   rÈ   r¼   s      r   Ú
transcribezRNNT.transcribeD  s   € ð6 ×Ñ ¨Ó8Ð8r   c                 ó*   — | j                  |||¬«      S )a  Applies prediction network to targets.

        B: batch size;
        U: maximum target sequence length in batch;
        D: feature dimension of each target sequence frame.

        Args:
            targets (torch.Tensor): target sequences, with shape `(B, U)` and each element
                mapping to a target symbol, i.e. in range `[0, num_symbols)`.
            target_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``targets``.
            state (List[List[torch.Tensor]] or None): list of lists of tensors
                representing internal state generated in preceding invocation
                of ``predict``.

        Returns:
            (torch.Tensor, torch.Tensor, List[List[torch.Tensor]]):
                torch.Tensor
                    output frame sequences, with shape `(B, U, output_dim)`.
                torch.Tensor
                    output lengths, with shape `(B,)` and i-th element representing
                    number of valid elements for i-th batch element in output.
                List[List[torch.Tensor]]
                    output states; list of lists of tensors
                    representing internal state generated in current invocation of ``predict``.
        rÌ   )rÅ   )r   rÉ   r¾   rD   s       r   ÚpredictzRNNT.predicta  s   € ðB �~‰~ G°^È5ˆ~ÓQÐQr   r»   r½   c                 ó>   — | j                  ||||¬«      \  }}}|||fS )a¶  Applies joint network to source and target encodings.

        B: batch size;
        T: maximum source sequence length in batch;
        U: maximum target sequence length in batch;
        D: dimension of each source and target sequence encoding.

        Args:
            source_encodings (torch.Tensor): source encoding sequences, with
                shape `(B, T, D)`.
            source_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                valid sequence length of i-th batch element in ``source_encodings``.
            target_encodings (torch.Tensor): target encoding sequences, with shape `(B, U, D)`.
            target_lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                valid sequence length of i-th batch element in ``target_encodings``.

        Returns:
            (torch.Tensor, torch.Tensor, torch.Tensor):
                torch.Tensor
                    joint network output, with shape `(B, T, U, output_dim)`.
                torch.Tensor
                    output source lengths, with shape `(B,)` and i-th element representing
                    number of valid elements along dim 1 for i-th batch element in joint network output.
                torch.Tensor
                    output target lengths, with shape `(B,)` and i-th element representing
                    number of valid elements along dim 2 for i-th batch element in joint network output.
        rÍ   )rÆ   )r   r»   r¼   r½   r¾   r&   s         r   Újoinz	RNNT.join„  s:   € ðF 26·±Ø-Ø)Ø-Ø)ð	 2=ó 2
Ñ.ˆ� ð �~ ~Ð5Ð5r   r   )r(   r)   r*   r+   r`   r“   r¶   r   r-   r.   r   r   r   r'   r�   r‘   rÏ   rÑ   rÓ   rÕ   r/   r0   s   @r   r	   r	   Å  s�  ø„ ñð  Lð ¸Zð ÐQXð Ð]aõ ð ?CñA
à—‘ðA
ð Ÿ™ðA
ð —‘ð	A
ð
 Ÿ™ðA
ð " $ t¨E¯L©LÑ'9Ñ":Ñ;ðA
ð 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¸dÀ5Ç<Á<Ñ>PÑ9QÐQÑ	RóA
ðF ‡Y�Y×Ñð"Fà—‘ð"Fð Ÿ™ð"Fð ˜˜T %§,¡,Ñ/Ñ0Ñ1ð	"Fð
 
ˆu�|‰|˜UŸ\™\¨4°°U·\±\Ñ0BÑ+CÐCÑ	Dò"Fó ð"FðH ‡Y�Y×Ñð9à—‘ð9ð Ÿ™ð9ð 
ˆu�|‰|˜UŸ\™\Ð)Ñ	*ò	9ó ð9ð8 ‡Y�Y×Ñð Rà—‘ð Rð Ÿ™ð Rð ˜˜T %§,¡,Ñ/Ñ0Ñ1ð	 Rð
 
ˆu�|‰|˜UŸ\™\¨4°°U·\±\Ñ0BÑ+CÐCÑ	Dò Ró ð RðD ‡Y�Y×Ñð(6àŸ,™,ð(6ð Ÿ™ð(6ð  Ÿ,™,ð	(6ð
 Ÿ™ð(6ð 
ˆu�|‰|˜UŸ\™\¨5¯<©<Ð7Ñ	8ò(6ó ô(6r   r	   r3   Úencoding_dimr”   rq   rr   rs   rt   ru   rv   rw   rk   rl   rx   rm   rn   ro   r•   r–   r˜   r™   rš   r   c                 óŽ   — t        | ||||||||	|
|||||¬«      }t        ||||||||¬«      }t        ||«      }t        |||«      S )a 
  Builds Emformer-based :class:`~torchaudio.models.RNNT`.

    Note:
        For non-streaming inference, the expectation is for `transcribe` to be called on input
        sequences right-concatenated with `right_context_length` frames.

        For streaming inference, the expectation is for `transcribe_streaming` to be called
        on input chunks comprising `segment_length` frames right-concatenated with `right_context_length`
        frames.

    Args:
        input_dim (int): dimension of input sequence frames passed to transcription network.
        encoding_dim (int): dimension of transcription- and prediction-network-generated encodings
            passed to joint network.
        num_symbols (int): cardinality of set of target tokens.
        segment_length (int): length of input segment expressed as number of frames.
        right_context_length (int): length of right context expressed as number of frames.
        time_reduction_input_dim (int): dimension to scale each element in input sequences to
            prior to applying time reduction block.
        time_reduction_stride (int): factor by which to reduce length of input sequence.
        transformer_num_heads (int): number of attention heads in each Emformer layer.
        transformer_ffn_dim (int): hidden layer dimension of each Emformer layer's feedforward network.
        transformer_num_layers (int): number of Emformer layers to instantiate.
        transformer_left_context_length (int): length of left context considered by Emformer.
        transformer_dropout (float): Emformer dropout probability.
        transformer_activation (str): activation function to use in each Emformer layer's
            feedforward network. Must be one of ("relu", "gelu", "silu").
        transformer_max_memory_size (int): maximum number of memory elements to use.
        transformer_weight_init_scale_strategy (str): per-layer weight initialization scaling
            strategy. Must be one of ("depthwise", "constant", ``None``).
        transformer_tanh_on_mem (bool): if ``True``, applies tanh to memory elements.
        symbol_embedding_dim (int): dimension of each target token embedding.
        num_lstm_layers (int): number of LSTM layers to instantiate.
        lstm_layer_norm (bool): if ``True``, enables layer normalization for LSTM layers.
        lstm_layer_norm_epsilon (float): value of epsilon to use in LSTM layer normalization layers.
        lstm_dropout (float): LSTM dropout probability.

    Returns:
        RNNT:
            Emformer RNN-T model.
    )r3   rp   rq   rr   rs   rt   ru   rv   rw   rk   rl   rx   rm   rn   ro   )r•   r–   r—   r˜   r™   rš   )rg   r“   r¶   r	   )r3   rÖ   r”   rq   rr   rs   rt   ru   rv   rw   rk   rl   rx   rm   rn   ro   r•   r–   r˜   r™   rš   ÚencoderrÅ   rÆ   s                           r   r   r   °  s}   € ôB ØØØ%Ø1Ø!9Ø3Ø3Ø/Ø5Ø/Ø5Ø(GØ$?Ø/UØ 7ô€Gô" ØØØ1Ø'Ø,Ø'Ø 7Ø!ô	€Iô �\ ;Ó/€FÜ�˜ FÓ+Ð+r   c                 ó’   — t        d(i dd“dd“d| “dd“dd	“d
d“dd	“dd“dd“dd“dd“dd“dd“dd“dd“dd“dd “d!d"“d#d“d$d%“d&d'“ŽS ))zÓBuilds basic version of Emformer-based :class:`~torchaudio.models.RNNT`.

    Args:
        num_symbols (int): The size of target token lexicon.

    Returns:
        RNNT:
            Emformer RNN-T model.
    r3   éP   rÖ   i   r”   rq   é   rr   r8   rs   é€   rt   ru   é   rv   i   rw   é   rk   gš™™™™™¹?rl   Úgelurx   é   rm   r   rn   rj   ro   Tr•   i   r–   é   r˜   r™   gü©ñÒMbP?rš   g333333Ó?rb   )r   )r”   s    r   r
   r
     sß   € ô ò Ùðáðñ  ðñ ð	ñ
 ðñ "%ðñ  ðñ  ðñ !ðñ  "ðñ  ðñ  &ðñ )+ðñ %&ðñ 0;ðñ  !%ð!ñ" !ð#ñ$ ð%ñ& ð'ñ( !%ð)ñ* ð+ð r   )Úabcr   r   Útypingr   r   r   r-   Útorchaudio.modelsr   Ú__all__r<   ÚModuler   r2   r`   rg   r“   r¶   r	   r,   r^   r�   r]   r   r
   rb   r   r   Ú<module>rç      s¾  ðß #ß (Ñ (ã Ý &ò @€ô)�U—X‘X—_‘_ô )ôXO�%—(‘(—/‘/ô Oôd�3ô ôMG�u—x‘x—‘¨ô MGô`]J�—‘—‘ô ]Jô@:6ˆe�h‰h�o‰oô :6ôzh6ˆ5�8‰8�?‰?ô h6ðV],àð],ð ð],ð ð	],ð
 ð],ð ð],ð "ð],ð ð],ð ð],ð ð],ð  ð],ð ð],ð  ð],ð &)ð],ð "%ð],ð  -0ð!],ð" "ð#],ð$ ð%],ð& ð'],ð( ð)],ð* #ð+],ð, ð-],ð. 
ó/],ð@  Cð  ¨Dô  r   