Ë
    óÿæi†“  ã                   ó¤  — d dl Z d dlmZmZmZ d dlZdgZdej                  dej                  fd„Z	 ddej                  dej                  d	ej                  dej                  d
ej                  deej                     deej                     fd„Z	de
dej                  j                  fd„Zdee
   dedeee      fd„Zdee   dee   dedej$                  dej                  f
d„Z G d„ dej                  j                  «      Z G d„ dej                  j                  «      Z G d„ dej                  j                  «      Z G d„ de«      Zy)é    N)ÚListÚOptionalÚTupleÚEmformerÚlengthsÚreturnc                 ó  — | j                   d   }t        t        j                  | «      j	                  «       «      }t        j
                  || j                  | j                  ¬«      j                  ||«      | j                  d«      k\  }|S )Nr   )ÚdeviceÚdtypeé   )
ÚshapeÚintÚtorchÚmaxÚitemÚaranger
   r   ÚexpandÚ	unsqueeze)r   Ú
batch_sizeÚ
max_lengthÚpadding_masks       úo/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchaudio/models/emformer.pyÚ_lengths_to_padding_maskr   
   sq   € Ø—‘˜qÑ!€JÜ”U—Y‘Y˜wÓ'×,Ñ,Ó.Ó/€JÜ—<‘< 
°7·>±>ÈÏÉÔW×^Ñ^Ø�Jóà	×	Ñ	˜1Ó	ñ€Lð Ðó    Ú	utteranceÚright_contextÚsummaryÚmemsÚleft_context_keyc                 óŠ  — |j                  d«      | j                  d«      z   |j                  d«      z   }|j                  d«      }|dk(  rd }|S |t        j                  |«      j                  «       z
  |j                  d«      z
  }	|�|j                  d«      nd}
||j                  d«      z   |	z   |
z   }t	        |¬«      }|S )Nr   r   )r   )Úsizer   r   r   r   )r   r   r   r   r   r   ÚTÚBr   Úright_context_blocks_lengthÚleft_context_blocks_lengthÚklengthss               r   Ú_gen_padding_maskr'      sÈ   € ð 	×Ñ˜1Ó 	§¡¨qÓ 1Ñ1°G·L±LÀ³OÑC€AØ×Ñ˜1Ó€AØˆA‚vØˆð Ðð	 '(¬%¯)©)°GÓ*<×*@Ñ*@Ó*BÑ&BÀWÇ\Á\ÐRSÃ_Ñ&TÐ#ØAQÐA]Ð%5×%:Ñ%:¸1Ô%=ÐcdÐ"Ø˜TŸY™Y q›\Ñ)Ð,GÑGÐJdÑdˆÜ/¸ÔAˆØÐr   Ú
activationc                 óð   — | dk(  rt         j                  j                  «       S | dk(  rt         j                  j                  «       S | dk(  rt         j                  j	                  «       S t        d| › �«      ‚)NÚreluÚgeluÚsiluzUnsupported activation )r   ÚnnÚReLUÚGELUÚSiLUÚ
ValueError)r(   s    r   Ú_get_activation_moduler2   '   s]   € Ø�VÒÜ�x‰x�}‰}‹ÐØ	�vÒ	Ü�x‰x�}‰}‹ÐØ	�vÒ	Ü�x‰x�}‰}‹ÐäÐ2°:°,Ð?Ó@Ð@r   Úweight_init_scale_strategyÚ
num_layersc                 óH  — | €t        |«      D �cg c]  }d ‘Œ c}S | dk(  r2t        |«      D �cg c]  }dt        j                  |dz   «      z  ‘Œ c}S | dk(  r/t        |«      D �cg c]  }dt        j                  d«      z  ‘Œ c}S t        d| › �«      ‚c c}w c c}w c c}w )NÚ	depthwiseg      ð?r   Úconstanté   z-Unsupported weight_init_scale_strategy value )ÚrangeÚmathÚsqrtr1   )r3   r4   Ú_Ú	layer_idxs       r   Ú_get_weight_init_gainsr>   2   s«   € Ø!Ð)Ü# JÔ/Ó0Ñ/˜’Ð/Ñ0Ð0Ø	# {Ò	2Ü@EÀjÔ@QÓRÑ@Q°9�”d—i‘i 	¨A¡Ó.Ó.Ð@QÑRÐRØ	# zÒ	1Ü49¸*Ô4EÓFÑ4E y�”d—i‘i “lÓ"Ð4EÑFÐFäÐHÐIcÐHdÐeÓfÐfùò 1ùâRùâFs   �	B®"BÁ%BÚ
col_widthsÚcol_maskÚnum_rowsr
   c           	      ó  — t        | «      t        |«      k7  rt        d«      ‚t        | |«      D ��cg c]7  \  }}|rt        j                  |||¬«      nt        j
                  |||¬«      ‘Œ9 }}}t        j                  |d¬«      S c c}}w )Nz0Length of col_widths must match that of col_mask©r
   r   ©Údim)Úlenr1   Úzipr   ÚonesÚzerosÚcat)r?   r@   rA   r
   Ú	col_widthÚis_ones_colÚ
mask_blocks          r   Ú_gen_attention_mask_blockrN   =   s“   € ô ˆ:ƒœ#˜h›-Ò'ÜÐKÓLÐLô '*¨*°hÔ&?ô	ñ '@Ñ"ˆI�{ñ ô 	�
‰
�8˜Y¨vÕ6ä�[‰[˜ 9°VÔ<ñ	=ð '@ð	 ñ ô �9‰9�Z QÔ'Ð'ùós   ²<Bc                   óv  ‡ — e Zd ZdZ	 	 	 	 ddedededee   dedefˆ fd„Zd	e	j                  d
e	j                  dee	j                  e	j                  f   fd„Zde	j                  de	j                  dee	j                     de	j                  fd„Z	 	 dde	j                  de	j                  de	j                  de	j                  d
e	j                  de	j                  dee	j                     dee	j                     dee	j                  e	j                  e	j                  e	j                  f   fd„Zde	j                  de	j                  de	j                  de	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	j                  de	j                  d
e	j                  de	j                  de	j                  dee	j                  e	j                  e	j                  e	j                  f   fd„«       Zˆ xZS )Ú_EmformerAttentiona_  Emformer layer attention module.

    Args:
        input_dim (int): input dimension.
        num_heads (int): number of attention heads in each Emformer layer.
        dropout (float, optional): dropout probability. (Default: 0.0)
        weight_init_gain (float or None, optional): scale factor to apply when initializing
            attention module parameters. (Default: ``None``)
        tanh_on_mem (bool, optional): if ``True``, applies tanh to memory elements. (Default: ``False``)
        negative_inf (float, optional): value to use for negative infinity in attention weights. (Default: -1e8)
    Ú	input_dimÚ	num_headsÚdropoutÚweight_init_gainÚtanh_on_memÚnegative_infc                 óÐ  •— t         ‰| �  «        ||z  dk7  rt        d|› d|› d�«      ‚|| _        || _        || _        || _        || _        | j                  | j                  z  dz  | _        t        j                  j                  |d|z  d¬«      | _        t        j                  j                  ||d¬«      | _        t        j                  j                  ||d¬«      | _        |rt        j                  j                  j!                  | j                  j"                  |¬	«       t        j                  j                  j!                  | j                  j"                  |¬	«       y y )
Nr   zinput_dim (z") is not a multiple of num_heads (z).g      à¿r8   T)Úbias)Úgain)ÚsuperÚ__init__r1   rQ   rR   rS   rU   rV   Úscalingr   r-   ÚLinearÚemb_to_key_valueÚemb_to_queryÚout_projÚinitÚxavier_uniform_Úweight)ÚselfrQ   rR   rS   rT   rU   rV   Ú	__class__s          €r   r[   z_EmformerAttention.__init__Y   s%  ø€ ô 	‰ÑÔà�yÑ  AÒ%Ü˜{¨9¨+Ð5WÐXaÐWbÐbdÐeÓfÐfà"ˆŒØ"ˆŒØˆŒØ&ˆÔØ(ˆÔàŸ™¨$¯.©.Ñ8¸TÑAˆŒä %§¡§¡°	¸1¸y¹=Èt Ó TˆÔÜ!ŸH™HŸO™O¨I°yÀt˜OÓLˆÔÜŸ™Ÿ™¨	°9À4˜ÓHˆŒáÜ�H‰H�M‰M×)Ñ)¨$×*?Ñ*?×*FÑ*FÐM]Ð)Ô^Ü�H‰H�M‰M×)Ñ)¨$×*;Ñ*;×*BÑ*BÐIYÐ)ÕZð r   Úinputr   r   c                 óÚ   — |j                   \  }}}|j                  d«      dz   }|d ||z
   }t        j                  ||g«      }| j	                  |«      j                  dd¬«      \  }}	||	fS )Nr   r   r8   ©ÚchunksrE   )r   r!   r   rJ   r^   Úchunk)
rd   rf   r   r"   r<   Úsummary_lengthÚright_ctx_utterance_blockÚmems_right_ctx_utterance_blockÚkeyÚvalues
             r   Ú_gen_key_valuez!_EmformerAttention._gen_key_valuew   s|   € Ø—+‘+‰ˆˆ1ˆaØŸ™ 1›¨Ñ)ˆØ$)Ð*>¨A°Ñ,>Ð$?Ð!Ü).¯©°DÐ:SÐ3TÓ)UÐ&Ø×*Ñ*Ð+IÓJ×PÑPÐXYÐ_`ÐPÓa‰
ˆˆUØ�EˆzÐr   Úattention_weightsÚattention_maskr   c                 ó
  — |j                  «       }|j                  |j                  d«      | j                  «      }|j	                  d«      }|j	                  d«      | j
                  z  }|�•|j                  || j
                  |d«      }|j                  |j                  d«      j                  d«      j                  t        j                  «      | j                  «      }|j                  || j
                  z  |d«      }t        j                  j                  j                  |d¬«      j                  |«      }t        j                  j                  j                  |t        | j                  «      | j                  ¬«      S )Nr   r   éÿÿÿÿr8   rD   )ÚpÚtraining)ÚfloatÚmasked_fillr   rV   r!   rR   ÚviewÚtor   Úboolr-   Ú
functionalÚsoftmaxÚtype_asrS   rv   )rd   rq   rr   r   Úattention_weights_floatr"   r#   Úattention_probss           r   Ú_gen_attention_probsz'_EmformerAttention._gen_attention_probs   sG  € ð #4×"9Ñ"9Ó";ÐØ"9×"EÑ"EÀn×F^ÑF^Ð_`ÓFaÐcg×ctÑctÓ"uÐØ×"Ñ" 1Ó%ˆØ×"Ñ" 1Ó%¨¯©Ñ7ˆØÐ#Ø&=×&BÑ&BÀ1ÀdÇnÁnÐVWÐY[Ó&\Ð#Ø&=×&IÑ&IØ×&Ñ& qÓ)×3Ñ3°AÓ6×9Ñ9¼%¿*¹*ÓEÀt×GXÑGXó'Ð#ð '>×&BÑ&BÀ1ÀtÇ~Á~ÑCUÐWXÐZ\Ó&]Ð#ÜŸ(™(×-Ñ-×5Ñ5Ð6MÐSUÐ5ÓV×^Ñ^Ð_pÓqˆÜ�x‰x×"Ñ"×*Ñ*¨?¼eÀDÇLÁLÓ>QÐ\`×\iÑ\iÐ*ÓjÐjr   r   r   r   r   r   Úleft_context_valc	           	      ód  — |j                  d«      }	|j                  d«      |j                  d«      z   |j                  d«      z   }
| j                  t        j                  |||g«      «      }| j	                  t        j                  |||g«      «      j                  dd¬«      \  }}|�¾|�¼|
t        j                  |«      j                  «       z
  |j                  d«      z
  }t        j                  |d |j                  d«      |z    |||j                  d«      |z   d  g«      }t        j                  |d |j                  d«      |z    |||j                  d«      |z   d  g«      }|||fD �cg c]W  }|j                  «       j                  d|	| j                  z  | j                  | j                  z  «      j                  dd«      ‘ŒY c}\  }}}t        j                  || j                  z  |j                  dd«      «      }t        ||||||«      }| j!                  |||«      }t        j                  ||«      }|j"                  |	| j                  z  |
| j                  | j                  z  fk7  rt%        d«      ‚|j                  dd«      j                  «       j                  |
|	| j                  «      }| j'                  |«      }|j                  d«      }|d |
|z
   }||
|z
  d  }| j(                  rt        j*                  |«      }nt        j,                  |dd¬	«      }||||fS c c}w )
Nr   r   r8   rh   rt   z+Computed attention has incorrect dimensionsiöÿÿÿé
   )Úminr   )r!   r_   r   rJ   r^   rj   r   r   Ú
contiguousry   rR   rQ   Ú	transposeÚbmmr\   r'   r�   r   ÚAssertionErrorr`   rU   ÚtanhÚclamp)rd   r   r   r   r   r   rr   r   r‚   r#   r"   Úqueryrn   ro   r$   ÚtensorÚreshaped_queryÚreshaped_keyÚreshaped_valuerq   r   r€   Ú	attentionÚoutput_right_context_memsrk   Úoutput_right_contextÚoutput_memss                              r   Ú_forward_implz _EmformerAttention._forward_impl’   s$  € ð �N‰N˜1ÓˆØ×Ñ˜qÓ! I§N¡N°1Ó$5Ñ5¸¿¹ÀQ»ÑGˆð ×!Ñ!¤%§)¡)¨]¸IÀwÐ,OÓ"PÓQˆð ×*Ñ*¬5¯9©9°d¸MÈ9Ð5UÓ+VÓW×]Ñ]ÐefÐlmÐ]Ón‰
ˆˆUàÐ'Ð,<Ð,HØ*+¬e¯i©i¸Ó.@×.DÑ.DÓ.FÑ*FÈÏÉÐVWËÑ*XÐ'Ü—)‘)àÐD˜$Ÿ)™) A›,Ð)DÑDÐEØ$Ø˜Ÿ	™	 !›Ð'BÑBÐDÐEðóˆCô —I‘IàÐF˜DŸI™I a›LÐ+FÑFÐGØ$Ø˜$Ÿ)™) A›,Ð)DÑDÐFÐGðóˆEð ! # uÑ-ó8
á-�ð ×ÑÓ×$Ñ$ R¨¨T¯^©^Ñ);¸T¿^¹^ÈtÏ~É~Ñ=]Ó^×hÑhÐijÐlmÕnØ-ñ8
Ñ4ˆ˜ nô "ŸI™I n°t·|±|Ñ&CÀ\×E[ÑE[Ð\]Ð_`ÓEaÓbÐô )¨°MÀ7ÈGÐUYÐ[kÓlˆð ×3Ñ3Ð4EÀ~ÐWcÓdˆô —I‘I˜o¨~Ó>ˆ	Ø�?‰?Ø�—‘ÑØØ�N‰N˜dŸn™nÑ,ð
ò 
ô
 !Ð!NÓOÐOØ×'Ñ'¨¨1Ó-×8Ñ8Ó:×?Ñ?ÀÀ1ÀdÇnÁnÓUˆ	ð %)§M¡M°)Ó$<Ð!à Ÿ™ a›ˆØ8Ð9M¸1¸~Ñ;MÐNÐØ/°°NÑ0BÐ0DÐEˆØ×ÒÜŸ*™* [Ó1‰KäŸ+™+ k°sÀÔCˆKà# [°#°uÐ<Ð<ùòC8
s   Å0AL-c                 óF   — | j                  ||||||«      \  }}}	}	||dd fS )ac  Forward pass for training.

        B: batch size;
        D: feature dimension of each frame;
        T: number of utterance frames;
        R: number of right context frames;
        S: number of summary elements;
        M: number of memory elements.

        Args:
            utterance (torch.Tensor): utterance frames, with shape `(T, B, D)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``utterance``.
            right_context (torch.Tensor): right context frames, with shape `(R, B, D)`.
            summary (torch.Tensor): summary elements, with shape `(S, B, D)`.
            mems (torch.Tensor): memory elements, with shape `(M, B, D)`.
            attention_mask (torch.Tensor): attention mask for underlying attention module.

        Returns:
            (Tensor, Tensor):
                Tensor
                    output frames corresponding to utterance and right_context, with shape `(T + R, B, D)`.
                Tensor
                    updated memory elements, with shape `(M, B, D)`.
        Nrt   )r•   )
rd   r   r   r   r   r   rr   Úoutputr”   r<   s
             r   Úforwardz_EmformerAttention.forwardÛ   s=   € ðD %)×$6Ñ$6°yÀ'È=ÐZaÐcgÐiwÓ$xÑ!ˆ�˜Q Ø�{ 3 BÐ'Ð'Ð'r   c           
      ó€  — |j                  d«      |j                  d«      z   |j                  d«      z   }|j                  d«      |j                  d«      z   |j                  d«      z   |j                  d«      z   }	t        j                  ||	«      j                  t        j                  |j
                  ¬«      }
d|
dd|j                  d«      …f<   | j                  ||||||
||¬«      \  }}}}||||j                  d«      |j                  d«      z   d ||j                  d«      |j                  d«      z   d fS )a½  Forward pass for inference.

        B: batch size;
        D: feature dimension of each frame;
        T: number of utterance frames;
        R: number of right context frames;
        S: number of summary elements;
        M: number of memory elements.

        Args:
            utterance (torch.Tensor): utterance frames, with shape `(T, B, D)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``utterance``.
            right_context (torch.Tensor): right context frames, with shape `(R, B, D)`.
            summary (torch.Tensor): summary elements, with shape `(S, B, D)`.
            mems (torch.Tensor): memory elements, with shape `(M, B, D)`.
            left_context_key (torch.Tensor): left context attention key computed from preceding invocation.
            left_context_val (torch.Tensor): left context attention value computed from preceding invocation.

        Returns:
            (Tensor, Tensor, Tensor, and Tensor):
                Tensor
                    output frames corresponding to utterance and right_context, with shape `(T + R, B, D)`.
                Tensor
                    updated memory elements, with shape `(M, B, D)`.
                Tensor
                    attention key computed for left context and utterance.
                Tensor
                    attention value computed for left context and utterance.
        r   ©r   r
   Trt   N)r   r‚   )r!   r   rI   rz   r{   r
   r•   )rd   r   r   r   r   r   r   r‚   Ú	query_dimÚkey_dimrr   r—   r”   rn   ro   s                  r   Úinferz_EmformerAttention.infer   sA  € ðR "×&Ñ& qÓ)¨I¯N©N¸1Ó,=Ñ=ÀÇÁÈQÃÑOˆ	Ø×$Ñ$ QÓ'¨)¯.©.¸Ó*;Ñ;¸d¿i¹iÈ»lÑJÐM]×MbÑMbÐcdÓMeÑeˆÜŸ™ Y°Ó8×;Ñ;Ä%Ç*Á*ÐU^×UeÑUeÐ;ÓfˆØ-1ˆ�r˜>˜TŸY™Y q›\˜>Ð)Ñ*Ø*.×*<Ñ*<ØØØØØØØ-Ø-ð +=ó 	+
Ñ'ˆ�˜S %ð ØØ�—	‘	˜!“˜}×1Ñ1°!Ó4Ñ4Ð6Ð7Ø�$—)‘)˜A“, ×!3Ñ!3°AÓ!6Ñ6Ð8Ð9ð	
ð 	
r   )ç        NFç    „×—Á)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   rw   r   r{   r[   r   ÚTensorr   rp   r�   r•   r˜   ÚjitÚexportr�   Ú__classcell__©re   s   @r   rP   rP   L   s®  ø„ ñ
ð  Ø,0Ø!Ø"ñ[àð[ð ð[ð ð	[ð
 # 5™/ð[ð ð[ð õ[ð< E§L¡Lð ¸¿¹ð ÈÈuÏ|É|Ð]b×]iÑ]iÐOiÑIjó ðkà Ÿ<™<ðkð Ÿ™ðkð ˜uŸ|™|Ñ,ð	kð
 
�‰ókð6 48Ø37ñG=à—<‘<ðG=ð —‘ðG=ð —|‘|ð	G=ð
 —‘ðG=ð �l‰lðG=ð Ÿ™ðG=ð # 5§<¡<Ñ0ðG=ð # 5§<¡<Ñ0ðG=ð 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¿¹ÐEÑ	FóG=ðR#(à—<‘<ð#(ð —‘ð#(ð —|‘|ð	#(ð
 —‘ð#(ð �l‰lð#(ð Ÿ™ð#(ð 
ˆu�|‰|˜UŸ\™\Ð)Ñ	*ó#(ðJ ‡Y�Y×Ñð;
à—<‘<ð;
ð —‘ð;
ð —|‘|ð	;
ð
 —‘ð;
ð �l‰lð;
ð  Ÿ,™,ð;
ð  Ÿ,™,ð;
ð 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¿¹ÐEÑ	Fò;
ó ô;
r   rP   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
ee   dedefˆ fd„Z	dedee
j                     dee
j                     fd„Zdee
j                     dee
j                  e
j                  e
j                  f   fd„Zde
j                  de
j                  dede
j                  dee
j                     dee
j                     fd„Zde
j                  de
j                  de
j                  de
j                  fd„Zde
j                  de
j                  dee
j                  e
j                  f   fd„Zde
j                  de
j                  de
j                  dee
j                  e
j                  f   fd„Zde
j                  de
j                  de
j                  de
j                  d ee
j                     dee
j                  e
j                  f   fd!„Zde
j                  de
j                  de
j                  de
j                  deee
j                        dee
j                  e
j                  ee
j                     f   fd"„Zde
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e
j0                  j2                  de
j                  de
j                  de
j                  deee
j                        de
j                  dee
j                  e
j                  ee
j                     e
j                  f   fd$„«       Zˆ xZS )&Ú_EmformerLayera$  Emformer layer that constitutes Emformer.

    Args:
        input_dim (int): input dimension.
        num_heads (int): number of attention heads.
        ffn_dim: (int): hidden layer dimension of feedforward network.
        segment_length (int): length of each input segment.
        dropout (float, optional): dropout probability. (Default: 0.0)
        activation (str, optional): activation function to use in feedforward network.
            Must be one of ("relu", "gelu", "silu"). (Default: "relu")
        left_context_length (int, optional): length of left context. (Default: 0)
        max_memory_size (int, optional): maximum number of memory elements to use. (Default: 0)
        weight_init_gain (float or None, optional): scale factor to apply when initializing
            attention module parameters. (Default: ``None``)
        tanh_on_mem (bool, optional): if ``True``, applies tanh to memory elements. (Default: ``False``)
        negative_inf (float, optional): value to use for negative infinity in attention weights. (Default: -1e8)
    rQ   rR   Úffn_dimÚsegment_lengthrS   r(   Úleft_context_lengthÚmax_memory_sizerT   rU   rV   c           
      óN  •— t         ‰| �  «        t        ||||	|
|¬«      | _        t        j
                  j                  |«      | _        t        j
                  j                  ||d¬«      | _	        t        |«      }t        j
                  j                  t        j
                  j                  |«      t        j
                  j                  ||«      |t        j
                  j                  |«      t        j
                  j                  ||«      t        j
                  j                  |«      «      | _        t        j
                  j                  |«      | _        t        j
                  j                  |«      | _        || _        || _        || _        || _        |dkD  | _        y )N)rQ   rR   rS   rT   rU   rV   T©Úkernel_sizeÚstrideÚ	ceil_moder   )rZ   r[   rP   r‘   r   r-   ÚDropoutrS   Ú	AvgPool1dÚ	memory_opr2   Ú
SequentialÚ	LayerNormr]   Úpos_ffÚlayer_norm_inputÚlayer_norm_outputr­   r¬   r®   rQ   Úuse_mem)rd   rQ   rR   r«   r¬   rS   r(   r­   r®   rT   rU   rV   Úactivation_modulere   s                €r   r[   z_EmformerLayer.__init__R  s=  ø€ ô 	‰ÑÔä+ØØØØ-Ø#Ø%ô
ˆŒô —x‘x×'Ñ'¨Ó0ˆŒÜŸ™×+Ñ+¸È~ÐimÐ+ÓnˆŒä2°:Ó>ÐÜ—h‘h×)Ñ)Ü�H‰H×Ñ˜yÓ)Ü�H‰H�O‰O˜I wÓ/ØÜ�H‰H×Ñ˜WÓ%Ü�H‰H�O‰O˜G YÓ/Ü�H‰H×Ñ˜WÓ%ó
ˆŒô !&§¡× 2Ñ 2°9Ó =ˆÔÜ!&§¡×!3Ñ!3°IÓ!>ˆÔà#6ˆÔ Ø,ˆÔØ.ˆÔØ"ˆŒà&¨Ñ*ˆ�r   r   r
   r   c                 ój  — t        j                  | j                  || j                  |¬«      }t        j                  | j                  || j                  |¬«      }t        j                  | j                  || j                  |¬«      }t        j                  d|t         j
                  |¬«      }||||gS )NrC   r   rš   )r   rI   r®   rQ   r­   Úint32)rd   r   r
   Úempty_memoryr   r‚   Úpast_lengths          r   Ú_init_statez_EmformerLayer._init_state€  sŠ   € Ü—{‘{ 4×#7Ñ#7¸ÀTÇ^Á^Ð\bÔcˆÜ Ÿ;™; t×'?Ñ'?ÀÈTÏ^É^ÐdjÔkÐÜ Ÿ;™; t×'?Ñ'?ÀÈTÏ^É^ÐdjÔkÐÜ—k‘k ! Z´u·{±{È6ÔRˆØÐ.Ð0@À+ÐNÐNr   Ústatec                 óT  — |d   d   d   j                  «       }t        | j                  |«      }t        | j                  t	        j
                  || j                  z  «      «      }|d   | j                  |z
  d  }|d   | j                  |z
  d  }|d   | j                  |z
  d  }|||fS )Né   r   r   r8   )r   r…   r­   r®   r:   Úceilr¬   )rd   rÃ   rÁ   Úpast_left_context_lengthÚpast_mem_lengthÚpre_memsÚlc_keyÚlc_vals           r   Ú_unpack_statez_EmformerLayer._unpack_state‡  s¸   € Ø˜A‘h˜q‘k !‘n×)Ñ)Ó+ˆÜ#& t×'?Ñ'?ÀÓ#MÐ Ü˜d×2Ñ2´D·I±I¸kÈD×L_ÑL_Ñ>_Ó4`ÓaˆØ˜‘8˜D×0Ñ0°?ÑBÐDÐEˆØ�q‘˜$×2Ñ2Ð5MÑMÐOÐPˆØ�q‘˜$×2Ñ2Ð5MÑMÐOÐPˆØ˜ Ð'Ð'r   Únext_kÚnext_vÚupdate_lengthr   c                 ób  — t        j                  |d   |g«      }t        j                  |d   |g«      }t        j                  |d   |g«      | j                   d  |d<   ||j                  d   | j                  z
  d  |d<   ||j                  d   | j                  z
  d  |d<   |d   |z   |d<   |S )Nr   r8   r   rÅ   )r   rJ   r®   r   r­   )rd   rÍ   rÎ   rÏ   r   rÃ   Únew_kÚnew_vs           r   Ú_pack_statez_EmformerLayer._pack_state�  s½   € ô —	‘	˜5 ™8 VÐ,Ó-ˆÜ—	‘	˜5 ™8 VÐ,Ó-ˆÜ—9‘9˜e A™h¨Ð-Ó.°×0DÑ0DÐ/DÐ/FÐGˆˆa‰Ø˜Ÿ™ Q™¨$×*BÑ*BÑBÐDÐEˆˆa‰Ø˜Ÿ™ Q™¨$×*BÑ*BÑBÐDÐEˆˆa‰Ø˜‘8˜mÑ+ˆˆa‰Øˆr   Ú	rc_outputr   r   c                 ó¢   — | j                  |«      t        j                  ||g«      z   }| j                  |«      |z   }| j	                  |«      }|S ©N)rS   r   rJ   r¹   r»   )rd   rÔ   r   r   Úresults        r   Ú_process_attention_outputz(_EmformerLayer._process_attention_output   sM   € ð —‘˜iÓ(¬5¯9©9°mÀYÐ5OÓ+PÑPˆØ—‘˜VÓ$ vÑ-ˆØ×'Ñ'¨Ó/ˆØˆr   c                 óž   — | j                  t        j                  ||g«      «      }||j                  d«      d  |d |j                  d«       fS ©Nr   )rº   r   rJ   r!   )rd   r   r   rº   s       r   Ú_apply_pre_attention_layer_normz._EmformerLayer._apply_pre_attention_layer_norm«  sY   € ð  ×0Ñ0´·±¸MÈ9Ð;UÓ1VÓWÐà˜]×/Ñ/°Ó2Ð4Ð5ØÐ4˜}×1Ñ1°!Ó4Ð5ð
ð 	
r   c                 óx   — | j                  |||«      }||j                  d«      d  |d |j                  d«       fS rÚ   )rØ   r!   )rd   rÔ   r   r   s       r   Ú_apply_post_attention_ffnz(_EmformerLayer._apply_post_attention_ffn´  sJ   € ð ×2Ñ2°9¸iÈÓWˆ	Ø˜×+Ñ+¨AÓ.Ð0Ð1°9Ð=T¸}×?QÑ?QÐRSÓ?TÐ3UÐUÐUr   r   rr   c                 óL  — |€t        d«      ‚| j                  r4| j                  |j                  ddd«      «      j                  ddd«      }n:t	        j
                  d«      j                  |j                  |j                  ¬«      }| j                  ||||||¬«      \  }}||fS )Nz;attention_mask must be not None when for_inference is Falser   r8   r   rš   )r   r   r   r   r   rr   )
r1   r¼   r¶   Úpermuter   Úemptyrz   r   r
   r‘   )	rd   r   r   r   r   rr   r   rÔ   Únext_ms	            r   Ú_apply_attention_forwardz'_EmformerLayer._apply_attention_forwardº  s§   € ð Ð!ÜÐZÓ[Ð[à�<Š<Ø—n‘n Y×%6Ñ%6°q¸!¸QÓ%?Ó@×HÑHÈÈAÈqÓQ‰Gä—k‘k !“n×'Ñ'¨i¯o©oÀi×FVÑFVÐ'ÓWˆGØ ŸN™NØØØ'ØØØ)ð +ó 
Ñˆ	�6ð ˜&Ð Ð r   c           	      ó&  — |€,| j                  |j                  d«      |j                  ¬«      }| j                  |«      \  }}}| j                  r9| j                  |j                  ddd«      «      j                  ddd«      }	|	d d }	n:t        j                  d«      j                  |j                  |j                  ¬«      }	| j                  j                  ||||	|||¬«      \  }
}}}| j                  |||j                  d«      ||«      }|
||fS )Nr   rC   r8   r   rš   )r   r   r   r   r   r   r‚   )rÂ   r!   r
   rÌ   r¼   r¶   rß   r   rà   rz   r   r‘   r�   rÓ   )rd   r   r   r   r   rÃ   rÉ   rÊ   rË   r   rÔ   rá   rÍ   rÎ   s                 r   Ú_apply_attention_inferz%_EmformerLayer._apply_attention_inferÓ  s  € ð ˆ=Ø×$Ñ$ Y§^¡^°AÓ%6¸y×?OÑ?OÐ$ÓPˆEØ#'×#5Ñ#5°eÓ#<Ñ ˆ�&˜&Ø�<Š<Ø—n‘n Y×%6Ñ%6°q¸!¸QÓ%?Ó@×HÑHÈÈAÈqÓQˆGØ˜b˜q�k‰Gä—k‘k !“n×'Ñ'¨i¯o©oÀi×FVÑFVÐ'ÓWˆGØ,0¯N©N×,@Ñ,@ØØØ'ØØØ#Ø#ð -Aó -
Ñ)ˆ	�6˜6 6ð × Ñ  ¨°·±ÀÓ1BÀDÈ%ÓPˆØ˜& %Ð'Ð'r   c                 ó’   — | j                  ||«      \  }}| j                  |||||«      \  }}	| j                  |||«      \  }
}|
||	fS )a1  Forward pass for training.

        B: batch size;
        D: feature dimension of each frame;
        T: number of utterance frames;
        R: number of right context frames;
        M: number of memory elements.

        Args:
            utterance (torch.Tensor): utterance frames, with shape `(T, B, D)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``utterance``.
            right_context (torch.Tensor): right context frames, with shape `(R, B, D)`.
            mems (torch.Tensor): memory elements, with shape `(M, B, D)`.
            attention_mask (torch.Tensor): attention mask for underlying attention module.

        Returns:
            (Tensor, Tensor, Tensor):
                Tensor
                    encoded utterance frames, with shape `(T, B, D)`.
                Tensor
                    updated right context frames, with shape `(R, B, D)`.
                Tensor
                    updated memory elements, with shape `(M, B, D)`.
        )rÛ   râ   rÝ   )rd   r   r   r   r   rr   Úlayer_norm_utteranceÚlayer_norm_right_contextrÔ   r”   Úoutput_utterancer“   s               r   r˜   z_EmformerLayer.forwardï  sr   € ðH ×0Ñ0°¸MÓJñ	
Ø Ø$à!%×!>Ñ!>Ø ØØ$ØØó"
Ñˆ	�;ð 26×1OÑ1OÐPYÐ[dÐfsÓ1tÑ.ÐÐ.ØÐ!5°{ÐBÐBr   c                 ó–   — | j                  ||«      \  }}| j                  |||||«      \  }}	}
| j                  |||«      \  }}|||
|	fS )a2  Forward pass for inference.

        B: batch size;
        D: feature dimension of each frame;
        T: number of utterance frames;
        R: number of right context frames;
        M: number of memory elements.

        Args:
            utterance (torch.Tensor): utterance frames, with shape `(T, B, D)`.
            lengths (torch.Tensor): with shape `(B,)` and i-th element representing
                number of valid frames for i-th batch element in ``utterance``.
            right_context (torch.Tensor): right context frames, with shape `(R, B, D)`.
            state (List[torch.Tensor] or None): list of tensors representing layer internal state
                generated in preceding invocation of ``infer``.
            mems (torch.Tensor): memory elements, with shape `(M, B, D)`.

        Returns:
            (Tensor, Tensor, List[torch.Tensor], Tensor):
                Tensor
                    encoded utterance frames, with shape `(T, B, D)`.
                Tensor
                    updated right context frames, with shape `(R, B, D)`.
                List[Tensor]
                    list of tensors representing layer internal state
                    generated in current invocation of ``infer``.
                Tensor
                    updated memory elements, with shape `(M, B, D)`.
        )rÛ   rä   rÝ   )rd   r   r   r   rÃ   r   ræ   rç   rÔ   r”   Úoutput_staterè   r“   s                r   r�   z_EmformerLayer.infer  ss   € ðR ×0Ñ0°¸MÓJñ	
Ø Ø$à/3×/JÑ/JØ  'Ð+CÀTÈ5ó0
Ñ,ˆ	�; ð 26×1OÑ1OÐPYÐ[dÐfsÓ1tÑ.ÐÐ.ØÐ!5°|À[ÐPÐPr   )rž   r*   r   r   NFrŸ   )r    r¡   r¢   r£   r   rw   Ústrr   r{   r[   r   r
   r   r¤   rÂ   r   rÌ   rÓ   rØ   rÛ   rÝ   râ   rä   r˜   r¥   r¦   r�   r§   r¨   s   @r   rª   rª   ?  s  ø„ ñð0 Ø Ø#$Ø Ø,0Ø!Ø"ñ,+àð,+ð ð,+ð ð	,+ð
 ð,+ð ð,+ð ð,+ð !ð,+ð ð,+ð # 5™/ð,+ð ð,+ð õ,+ð\O cð O°8¸E¿L¹LÑ3Ið OÈdÐSX×S_ÑS_ÑN`ó Oð( 4¨¯©Ñ#5ð (¸%ÀÇÁÈeÏlÉlÐ\a×\hÑ\hÐ@hÑ:ió (ðà—‘ðð —‘ðð ð	ð
 �l‰lðð �E—L‘LÑ!ðð 
ˆe�l‰lÑ	óð 	à—<‘<ð	ð —<‘<ð	ð —|‘|ð		ð
 
�‰ó	ð
ØŸ™ð
Ø6;·l±lð
à	ˆu�|‰|˜UŸ\™\Ð)Ñ	*ó
ðVØŸ™ðVØ27·,±,ðVØOTÏ|É|ðVà	ˆu�|‰|˜UŸ\™\Ð)Ñ	*óVð!à—<‘<ð!ð —‘ð!ð —|‘|ð	!ð
 �l‰lð!ð ! §¡Ñ.ð!ð 
ˆu�|‰|˜UŸ\™\Ð)Ñ	*ó!ð2(à—<‘<ð(ð —‘ð(ð —|‘|ð	(ð
 �l‰lð(ð ˜˜UŸ\™\Ñ*Ñ+ð(ð 
ˆu�|‰|˜UŸ\™\¨4°·±Ñ+=Ð=Ñ	>ó(ð8-Cà—<‘<ð-Cð —‘ð-Cð —|‘|ð	-Cð
 �l‰lð-Cð Ÿ™ð-Cð 
ˆu�|‰|˜UŸ\™\¨5¯<©<Ð7Ñ	8ó-Cð^ ‡Y�Y×Ñð-Qà—<‘<ð-Qð —‘ð-Qð —|‘|ð	-Qð
 ˜˜UŸ\™\Ñ*Ñ+ð-Qð �l‰lð-Qð 
ˆu�|‰|˜UŸ\™\¨4°·±Ñ+=¸u¿|¹|ÐKÑ	Lò-Qó ô-Qr   rª   c                   óL  ‡ — e Zd Z	 	 	 ddej                  j
                  dedededef
ˆ fd„Zdej                  dej                  fd	„Z	d
edede
e   fd„Zdej                  dej                  fd„Zdej                  dej                  deej                  ej                  f   fd„Zej                  j                   	 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 )Ú_EmformerImplÚemformer_layersr¬   r­   Úright_context_lengthr®   c                 óÊ   •— t         ‰| �  «        |dkD  | _        t        j                  j                  ||d¬«      | _        || _        || _        || _	        || _
        || _        y )Nr   Tr°   )rZ   r[   r¼   r   r-   rµ   r¶   rî   r­   rï   r¬   r®   )rd   rî   r¬   r­   rï   r®   re   s         €r   r[   z_EmformerImpl.__init__P  sj   ø€ ô 	‰ÑÔà&¨Ñ*ˆŒÜŸ™×+Ñ+Ø&Ø!Øð ,ó 
ˆŒð
  /ˆÔØ#6ˆÔ Ø$8ˆÔ!Ø,ˆÔØ.ˆÕr   rf   r   c                 ó~  — |j                   d   }t        j                  || j                  z
  | j                  z  «      }g }t        |dz
  «      D ]7  }|dz   | j                  z  }|| j                  z   }|j                  ||| «       Œ9 |j                  ||| j                  z
  d  «       t        j                  |«      S ©Nr   r   )	r   r:   rÆ   rï   r¬   r9   Úappendr   rJ   )rd   rf   r"   Únum_segsÚright_context_blocksÚseg_idxÚstartÚends           r   Ú_gen_right_contextz _EmformerImpl._gen_right_contextf  s¸   € Ø�K‰K˜‰NˆÜ—9‘9˜a $×";Ñ";Ñ;¸t×?RÑ?RÑRÓSˆØ!ÐÜ˜X¨™\Ö*ˆGØ˜q‘[ D×$7Ñ$7Ñ7ˆEØ˜$×3Ñ3Ñ3ˆCØ ×'Ñ'¨¨e°CÐ(8Õ9ð +ð 	×#Ñ# E¨!¨d×.GÑ.GÑ*GÐ*IÐ$JÔKÜ�y‰yÐ-Ó.Ð.r   rö   Úutterance_lengthc           
      óÊ  — t        j                  || j                  z  «      }| j                  }| j                  }||z  }||z   }t        || j                  z  |z
  d«      }t        |dz   | j                  z  |«      }	| j                  |z  }
| j                  r:t        || j                  z
  d«      }|dz
  }|||z
  ||z
  |||
|z
  ||	|z
  ||	z
  g	}|S |||
|z
  ||	|z
  ||	z
  g}|S rò   )	r:   rÆ   r¬   rï   r­   r   r…   r¼   r®   )rd   rö   rú   rô   ÚrcÚlcÚrc_startÚrc_endÚ	seg_startÚseg_endÚ	rc_lengthÚm_startÚ
mem_lengthr?   s                 r   Ú_gen_attention_mask_col_widthsz,_EmformerImpl._gen_attention_mask_col_widthsq  s(  € Ü—9‘9Ð-°×0CÑ0CÑCÓDˆØ×&Ñ&ˆØ×%Ñ%ˆØ˜R‘<ˆØ˜B‘ˆÜ˜ $×"5Ñ"5Ñ5¸Ñ:¸AÓ>ˆ	Ü�w ‘{ d×&9Ñ&9Ñ9Ð;KÓLˆØ×-Ñ-°Ñ8ˆ	à�<Š<Ü˜' D×$8Ñ$8Ñ8¸!Ó<ˆGØ! A™ˆJàØ˜'Ñ!Ø˜WÑ$ØØØ˜FÑ"ØØ˜)Ñ#Ø  7Ñ*ð
ˆJð* Ðð ØØ˜FÑ"ØØ˜)Ñ#Ø  7Ñ*ðˆJð Ðr   c                 ó¼  — |j                  d«      }t        j                  || j                  z  «      }g }g }g }| j                  r<d}t        |«      D �cg c]  }|dv ‘Œ }	}t        |«      D �cg c]  }|dv ‘Œ }
}|||g}n"d}t        |«      D �cg c]  }|dv ‘Œ }	}d }
||g}t        |«      D ]À  }| j                  ||«      }t        ||	| j                  |j                  «      }|j                  |«       t        ||	t        | j                  ||| j                  z  z
  «      |j                  «      }|j                  |«       |
€Œ˜t        ||
d|j                  «      }|j                  |«       ŒÂ dt        j                  |D �cg c]  }t        j                  |«      ‘Œ c}«      z
  j                  t        j                  «      }|S c c}w c c}w c c}w c c}w )Nr   é	   )r   é   é   )r  r	  é   )r   r  r   )r!   r:   rÆ   r¬   r¼   r9   r  rN   rï   r
   ró   r…   r   rJ   rz   r{   )rd   rf   rú   rô   Úrc_maskÚ
query_maskÚsummary_maskÚnum_colsÚidxÚrc_q_cols_maskÚs_cols_maskÚmasks_to_concatrö   r?   Úrc_mask_blockÚquery_mask_blockÚsummary_mask_blockÚmaskrr   s                      r   Ú_gen_attention_maskz!_EmformerImpl._gen_attention_mask•  sÝ  € Ø Ÿ:™: a›=ÐÜ—9‘9Ð-°×0CÑ0CÑCÓDˆàˆØˆ
Øˆà�<Š<ØˆHä:?À¼/ÓJ¹/°3˜c YÒ.¸/ˆNÐJä49¸(´OÓD±O¨S˜3 &š=°OˆKÐDØ&¨
°LÐA‰OàˆHä7<¸X´ÓG±°˜c Všm°ˆNÐGØˆKØ&¨
Ð3ˆOä˜X–ˆGØ×<Ñ<¸WÐFVÓWˆJä5Ø˜N¨D×,EÑ,EÀuÇ|Á|óˆMð �N‰N˜=Ô)ä8ØØÜØ×'Ñ'Ø$ w°×1DÑ1DÑ'DÑDóð —‘ó Ðð ×ÑÐ.Ô/àÑ&Ü%>¸zÈ;ÐXYÐ[`×[gÑ[gÓ%hÐ"Ø×#Ñ#Ð$6Õ7ð+ 'ð. œeŸi™iÁ_Ó(UÁ_¸T¬¯©°4­À_Ñ(UÓVÑV×ZÑZÔ[`×[eÑ[eÓfˆØÐùòG KùâDùò
 Hùò6 )Vs   ÁG
Á/GÂGÆG
r   c                 ó  — |j                  ddd«      }| j                  |«      }|d|j                  d«      | j                  z
   }| j	                  |«      }| j
                  r6| j                  |j                  ddd«      «      j                  ddd«      dd n9t        j                  d«      j                  |j                  |j                  ¬«      }|}| j                  D ]  } ||||||«      \  }}}Œ |j                  ddd«      |fS )aG  Forward pass for training and non-streaming inference.

        B: batch size;
        T: max number of input frames in batch;
        D: feature dimension of each frame.

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

        Returns:
            (Tensor, Tensor):
                Tensor
                    output frames, with shape `(B, T, D)`.
                Tensor
                    output lengths, with shape `(B,)` and i-th element representing
                    number of valid frames for i-th batch element in output frames.
        r   r   r8   Nrt   rš   )rß   rù   r!   rï   r  r¼   r¶   r   rà   rz   r   r
   rî   )	rd   rf   r   r   r   rr   r   r—   Úlayers	            r   r˜   z_EmformerImpl.forwardÅ  s  € ð* —‘˜a  AÓ&ˆØ×/Ñ/°Ó6ˆØÐE˜EŸJ™J q›M¨D×,EÑ,EÑEÐFˆ	Ø×1Ñ1°)Ó<ˆð �|Š|ð �N‰N˜9×,Ñ,¨Q°°1Ó5Ó6×>Ñ>¸qÀ!ÀQÓGÈÈÑLä—‘˜Q“×"Ñ"¨¯©¸U¿\¹\Ð"ÓJð 	ð
 ˆØ×)Ô)ˆEÙ*/°¸ÀÐPTÐVdÓ*eÑ'ˆF�M¡4ð *à�~‰~˜a  AÓ&¨Ð/Ð/r   Ústatesc                 óJ  — |j                  d«      | j                  | j                  z   k7  r8t        d| j                  | j                  z   › d|j                  d«      › d�«      ‚|j	                  ddd«      }|j                  d«      | j                  z
  }||d }|d| }t        j                  || j                  z
  d¬«      }| j                  r3| j                  |j	                  ddd«      «      j	                  ddd«      n9t        j                  d«      j                  |j                  |j                  ¬	«      }|}	g }
t        | j                  «      D ]7  \  }}|j                  |	|||€dn||   |«      \  }	}}}|
j!                  |«       Œ9 |	j	                  ddd«      ||
fS )
a±  Forward pass for streaming inference.

        B: batch size;
        D: feature dimension of each frame.

        Args:
            input (torch.Tensor): utterance frames right-padded with right context frames, with
                shape `(B, segment_length + 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``.
            states (List[List[torch.Tensor]] or None, optional): list of lists of tensors
                representing internal state generated in preceding invocation of ``infer``. (Default: ``None``)

        Returns:
            (Tensor, Tensor, List[List[Tensor]]):
                Tensor
                    output frames, with shape `(B, segment_length, D)`.
                Tensor
                    output lengths, with shape `(B,)` and i-th element representing
                    number of valid frames for i-th batch element in output frames.
                List[List[Tensor]]
                    output states; list of lists of tensors representing internal state
                    generated in current invocation of ``infer``.
        r   zIPer configured segment_length and right_context_length, expected size of z# for dimension 1 of input, but got Ú.r   r8   N)r…   rš   )r!   r¬   rï   r1   rß   r   r‹   r¼   r¶   rà   rz   r   r
   Ú	enumeraterî   r�   ró   )rd   rf   r   r  Úright_context_start_idxr   r   Úoutput_lengthsr   r—   Úoutput_statesr=   r  rê   s                 r   r�   z_EmformerImpl.inferè  s­  € ð> �:‰:�a‹=˜D×/Ñ/°$×2KÑ2KÑKÒKÜð&Ø&*×&9Ñ&9¸D×<UÑ<UÑ&UÐ%Vð WØ"ŸZ™Z¨›]˜O¨1ð.óð ð
 —‘˜a  AÓ&ˆØ"'§*¡*¨Q£-°$×2KÑ2KÑ"KÐØÐ5Ð6Ð7ˆØÐ2Ð2Ð3ˆ	ÜŸ™ W¨t×/HÑ/HÑ%HÈaÔPˆð �|Š|ð �N‰N˜9×,Ñ,¨Q°°1Ó5Ó6×>Ñ>¸qÀ!ÀQÔGä—‘˜Q“×"Ñ"¨¯©¸U¿\¹\Ð"ÓJð 	ð
 ˆØ24ˆÜ )¨$×*>Ñ*>Ö ?ÑˆI�uØ8=¿¹ØØØØ˜‘¨F°9Ñ,=Øó9Ñ5ˆF�M <°ð × Ñ  Õ.ð !@ð �~‰~˜a  AÓ&¨¸ÐEÐEr   )r   r   r   rÖ   )r    r¡   r¢   r   r-   Ú
ModuleListr   r[   r¤   rù   r   r  r  r   r˜   r¥   r¦   r   r�   r§   r¨   s   @r   rí   rí   O  sj  ø„ ð
 $%Ø$%Ø ñ/àŸ™×,Ñ,ð/ð ð/ð !ð	/ð
 "ð/ð õ/ð,	/¨¯©ð 	/¸¿¹ó 	/ð"°cð "ÈSð "ÐUYÐZ]ÑU^ó "ðH.¨¯©ð .¸%¿,¹,ó .ð`!0˜UŸ\™\ð !0°E·L±Lð !0ÀUÈ5Ï<É<ÐY^×YeÑYeÐKeÑEfó !0ðF ‡Y�Y×Ñð
 6:ñ	:Fà�|‰|ð:Fð —‘ð:Fð ˜˜d 5§<¡<Ñ0Ñ1Ñ2ð	:Fð
 
ˆu�|‰|˜UŸ\™\¨4°°U·\±\Ñ0BÑ+CÐCÑ	Dò:Fó ô:Fr   rí   c                   óp   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 ddededededededed	ed
ededee   dedefˆ fd„Z	ˆ xZ
S )r   a_  Emformer architecture introduced in
    *Emformer: Efficient Memory Transformer Based Acoustic Model for Low Latency Streaming Speech Recognition*
    :cite:`shi2021emformer`.

    See Also:
        * :func:`~torchaudio.models.emformer_rnnt_model`,
          :func:`~torchaudio.models.emformer_rnnt_base`: factory functions.
        * :class:`torchaudio.pipelines.RNNTBundle`: ASR pipelines with pretrained model.

    Args:
        input_dim (int): input dimension.
        num_heads (int): number of attention heads in each Emformer layer.
        ffn_dim (int): hidden layer dimension of each Emformer layer's feedforward network.
        num_layers (int): number of Emformer layers to instantiate.
        segment_length (int): length of each input segment.
        dropout (float, optional): dropout probability. (Default: 0.0)
        activation (str, optional): activation function to use in each Emformer layer's
            feedforward network. Must be one of ("relu", "gelu", "silu"). (Default: "relu")
        left_context_length (int, optional): length of left context. (Default: 0)
        right_context_length (int, optional): length of right context. (Default: 0)
        max_memory_size (int, optional): maximum number of memory elements to use. (Default: 0)
        weight_init_scale_strategy (str or None, optional): per-layer weight initialization scaling
            strategy. Must be one of ("depthwise", "constant", ``None``). (Default: "depthwise")
        tanh_on_mem (bool, optional): if ``True``, applies tanh to memory elements. (Default: ``False``)
        negative_inf (float, optional): value to use for negative infinity in attention weights. (Default: -1e8)

    Examples:
        >>> emformer = Emformer(512, 8, 2048, 20, 4, right_context_length=1)
        >>> input = torch.rand(128, 400, 512)  # batch, num_frames, feature_dim
        >>> lengths = torch.randint(1, 200, (128,))  # batch
        >>> output, lengths = emformer(input, lengths)
        >>> input = torch.rand(128, 5, 512)
        >>> lengths = torch.ones(128) * 5
        >>> output, lengths, states = emformer.infer(input, lengths, None)
    rQ   rR   r«   r4   r¬   rS   r(   r­   rï   r®   r3   rU   rV   c                 óê   •— t        ||«      }t        j                  j                  t	        |«      D �cg c]  }t        ||||||||
||   ||¬«      ‘Œ c}«      }t        ‰| �  ||||	|
¬«       y c c}w )N)rS   r(   r­   r®   rT   rU   rV   )r­   rï   r®   )r>   r   r-   r!  r9   rª   rZ   r[   )rd   rQ   rR   r«   r4   r¬   rS   r(   r­   rï   r®   r3   rU   rV   Úweight_init_gainsr=   rî   re   s                    €r   r[   zEmformer.__init__K  s    ø€ ô  3Ð3MÈzÓZÐÜŸ(™(×-Ñ-ô "' zÔ!2óñ "3�Iô ØØØØ"Ø#Ø)Ø(;Ø$3Ø%6°yÑ%AØ +Ø!-öð "3ñó
ˆô$ 	‰ÑØØØ 3Ø!5Ø+ð 	õ 	
ùò#s   ´ A0)rž   r*   r   r   r   r6   FrŸ   )r    r¡   r¢   r£   r   rw   rë   r   r{   r[   r§   r¨   s   @r   r   r   &  s±   ø„ ñ"ðV Ø Ø#$Ø$%Ø Ø4?Ø!Ø"ñ)
àð)
ð ð)
ð ð	)
ð
 ð)
ð ð)
ð ð)
ð ð)
ð !ð)
ð "ð)
ð ð)
ð %-¨S¡Mð)
ð ð)
ð ÷)
ñ )
r   rÖ   )r:   Útypingr   r   r   r   Ú__all__r¤   r   r'   rë   r-   ÚModuler2   r   rw   r>   r{   r
   rN   rP   rª   rí   r   © r   r   Ú<module>r)     s‰  ðÛ ß (Ñ (ã ð ˆ,€ð e§l¡lð °u·|±|ó ð 04ñØ�|‰|ðà—<‘<ðð �\‰\ðð �\‰\ð	ð
 �,‰,ðð ˜uŸ|™|Ñ,ðð ˆe�l‰lÑóð(A sð A¨u¯x©x¯©ó Aðg°xÀ±}ð gÐRUð gÐZ^Ð_gÐhmÑ_nÑZoó gð(Ø�S‘	ð(Ø%)¨$¡Zð(Ø;>ð(ØHMÏÉð(à
‡\�\ó(ôp
˜Ÿ™Ÿ™ô p
ôfMQ�U—X‘X—_‘_ô MQô`TF�E—H‘H—O‘Oô TFônN
ˆ}õ N
r   