Ë
    T^(h8Ø ã                   ó&  — d Z ddlZddl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ZddlmZ ddlmZ ddlmZ dd	lmZ dd
lmZ ddlmZmZ ddlmZmZmZmZmZmZm Z  ddl!m"Z" ddl#m$Z$m%Z%m&Z&m'Z'm(Z(m)Z)m*Z*m+Z+m,Z, ddl-m.Z. dZ/dZ0 e*«       rddl1m2Z3  e«       rddlm4Z4  e+jj                  e6«      Z7dZ8dZ9dZ:g d¢Z;dZ<dZ=dZ>dZ?dZ@dZAddgZBdZCd ZDe G d!„ d"e$«      «       ZE	 	 dvd#eeFeFf   d$eGd%eFd&eej�                     d'eFd(e
j’                  fd)„ZJ	 dwd*ed+eFd,ee
j’                     fd-„ZK G d.„ d/ej˜                  «      ZM G d0„ d1ej˜                  «      ZN G d2„ d3ej˜                  «      ZO G d4„ d5ej˜                  «      ZP G d6„ d7ej˜                  «      ZQ G d8„ d9ej˜                  «      ZR G d:„ d;eR«      ZS G d<„ d=ej˜                  «      ZT G d>„ d?ej˜                  «      ZU G d@„ dAeU«      ZV G dB„ dCeU«      ZWeUeWeVdDœZX G dE„ dFej˜                  «      ZY G dG„ dHej˜                  «      ZZ G dI„ dJej˜                  «      Z[ G dK„ dLej˜                  «      Z\ G dM„ dNej˜                  «      Z] G dO„ dPej˜                  «      Z^ G dQ„ dRej˜                  «      Z_ G dS„ dTej˜                  «      Z` G dU„ dVej˜                  «      Za G dW„ dXe"«      ZbdYZcdZZd e&d[ec«       G d\„ d]eb«      «       Ze e&d^ec«       G d_„ d`eb«      «       Zf e&daec«       G db„ dceb«      «       Zg e&ddecde«       G df„ dgeb«      «       Zh e&dhec«       G di„ djeb«      «       Zi e&dkec«       G dl„ dmeb«      «       Zj G dn„ doej˜                  «      Zk G dp„ dqej˜                  «      Zl e&drec«       G ds„ dteb«      «       Zmg du¢Zny)xzPyTorch Wav2Vec2 model.é    N)Ú	dataclass)ÚOptionalÚTupleÚUnion)Únn)ÚCrossEntropyLossé   )ÚACT2FN)Úis_deepspeed_zero3_enabled)Úis_fsdp_managed_module)Ú!flash_attn_supports_top_left_maskÚis_flash_attn_available)ÚBaseModelOutputÚCausalLMOutputÚMaskedLMOutputÚSequenceClassifierOutputÚTokenClassifierOutputÚWav2Vec2BaseModelOutputÚXVectorOutput)ÚPreTrainedModel)	ÚModelOutputÚadd_code_sample_docstringsÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚcached_fileÚis_peft_availableÚis_safetensors_availableÚloggingÚreplace_return_docstringsé   )ÚWav2Vec2Configzadapter.{}.binzadapter.{}.safetensors)Ú	load_file)Ú_flash_attention_forwardé   r!   zfacebook/wav2vec2-base-960h)r    i$  i   z['MISTER QUILTER IS THE APOSTLE OF THE MIDDLE CLASSES AND WE ARE GLAD TO WELCOME HIS GOSPEL'g=
×£p½J@zsuperb/wav2vec2-base-superb-ksz'_unknown_'g)\�Âõ(@zanton-l/wav2vec2-base-superb-sdzanton-l/wav2vec2-base-superb-svg\�Âõ(\ï?c                   ó^  — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eej                     ed<   dZeej                     ed<   dZeeej                        ed<   dZeeej                        ed<   dZeej                     ed	<   dZeej                     ed
<   y)ÚWav2Vec2ForPreTrainingOutputa1	  
    Output type of [`Wav2Vec2ForPreTraining`], with potential hidden states and attentions.

    Args:
        loss (*optional*, returned when `sample_negative_indices` are passed, `torch.FloatTensor` of shape `(1,)`):
            Total loss as the sum of the contrastive loss (L_m) and the diversity loss (L_d) as stated in the [official
            paper](https://arxiv.org/pdf/2006.11477.pdf) . (classification) loss.
        projected_states (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.proj_codevector_dim)`):
            Hidden-states of the model projected to *config.proj_codevector_dim* that can be used to predict the masked
            projected quantized states.
        projected_quantized_states (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.proj_codevector_dim)`):
            Quantized extracted feature vectors projected to *config.proj_codevector_dim* representing the positive
            target vectors for contrastive loss.
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
            shape `(batch_size, sequence_length, hidden_size)`.

            Hidden-states of the model at the output of each layer plus the initial embedding outputs.
        attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
            sequence_length)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
        contrastive_loss (*optional*, returned when `sample_negative_indices` are passed, `torch.FloatTensor` of shape `(1,)`):
            The contrastive loss (L_m) as stated in the [official paper](https://arxiv.org/pdf/2006.11477.pdf) .
        diversity_loss (*optional*, returned when `sample_negative_indices` are passed, `torch.FloatTensor` of shape `(1,)`):
            The diversity loss (L_d) as stated in the [official paper](https://arxiv.org/pdf/2006.11477.pdf) .
    NÚlossÚprojected_statesÚprojected_quantized_statesÚcodevector_perplexityÚhidden_statesÚ
attentionsÚcontrastive_lossÚdiversity_loss)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r'   r   ÚtorchÚFloatTensorÚ__annotations__r(   r)   r*   r+   r   r,   r-   r.   © ó    úl/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/wav2vec2/modeling_wav2vec2.pyr&   r&   a   s¿   … ñð< )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø48Ð�h˜u×0Ñ0Ñ1Ó8Ø>BÐ ¨×):Ñ):Ñ ;ÓBØ9=Ð˜8 E×$5Ñ$5Ñ6Ó=Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ó9Ø48Ð�h˜u×0Ñ0Ñ1Ó8Ø26€N�H˜U×.Ñ.Ñ/Ô6r7   r&   ÚshapeÚ	mask_probÚmask_lengthÚattention_maskÚ	min_masksÚreturnc                 óà  ‡‡‡‡‡— | \  }Š‰dk  rt        d«      ‚‰‰kD  rt        d‰› d‰› d�«      ‚t        j                  j                  d«      j	                  «       Šˆˆˆˆˆfd„}|�-|j                  «       j                  d«      j                  «       nt        |«      D �cg c]  }‰‘Œ c}}t        j                  |‰ft        ¬	«      }	g }
 |‰«      }|d
k(  r|	S |D ]¯  } ||«      }t        j                  j                  t        j                  |‰dz
  z
  «      |d¬«      }t        |«      d
k(  r‰dz
  }n|d
   }t        j                  |t        j                  ||z
  t        j                   ¬	«      |z  g«      }|
j#                  |«       Œ± t        j$                  |
«      }
t        j&                  |
dd…dd…df   ||‰f«      }
|
j)                  ||‰z  «      }
t        j                  ‰«      dddd…f   }t        j&                  |||‰f«      j)                  ||‰z  «      }|
|z   }
|
j+                  «       ‰dz
  kD  r‰dz
  |
|
‰dz
  kD  <   t        j,                  |	|
dd«       |	S c c}w )af  
    Computes random mask spans for a given shape. Used to implement [SpecAugment: A Simple Data Augmentation Method for
    ASR](https://arxiv.org/abs/1904.08779). Note that this method is not optimized to run on TPU and should be run on
    CPU as part of the preprocessing during training.

    Args:
        shape: The shape for which to compute masks. This should be of a tuple of size 2 where
               the first element is the batch size and the second element is the length of the axis to span.
        mask_prob:  The percentage of the whole axis (between 0 and 1) which will be masked. The number of
                    independently generated mask spans of length `mask_length` is computed by
                    `mask_prob*shape[1]/mask_length`. Note that due to overlaps, `mask_prob` is an upper bound and the
                    actual percentage will be smaller.
        mask_length: size of the mask
        min_masks: minimum number of masked spans
        attention_mask: A (right-padded) attention mask which independently shortens the feature axis of
                        each batch dimension.
    r    z&`mask_length` has to be bigger than 0.zO`mask_length` has to be smaller than `sequence_length`, but got `mask_length`: z and `sequence_length`: ú`c                 óœ   •— t        ‰| z  ‰z  ‰z   «      }t        |‰«      }|‰z  ‰kD  r‰‰z  }| ‰dz
  z
  |k  rt        | ‰dz
  z
  d«      }|S )z;Given input length, compute how many spans should be maskedr    r   )ÚintÚmax)Úinput_lengthÚnum_masked_spanÚepsilonr;   r:   r=   Úsequence_lengths     €€€€€r8   Úcompute_num_masked_spanz6_compute_mask_indices.<locals>.compute_num_masked_span±   so   ø€ ä˜i¨,Ñ6¸ÑDÀwÑNÓOˆÜ˜o¨yÓ9ˆð ˜[Ñ(¨?Ò:Ø-°Ñ<ˆOð ˜;¨™?Ñ+¨oÒ=Ü! ,°+À±/Ñ"BÀAÓFˆOàÐr7   Néÿÿÿÿ©Údtyper   F)Úreplace)Ú
ValueErrorÚnpÚrandomÚrandÚitemÚdetachÚsumÚtolistÚrangeÚzerosÚboolÚchoiceÚarangeÚlenÚconcatenateÚonesÚint32ÚappendÚarrayÚbroadcast_toÚreshaperC   Úput_along_axis)r9   r:   r;   r<   r=   Ú
batch_sizerH   Ú_Úinput_lengthsÚspec_aug_maskÚspec_aug_mask_idxsÚmax_num_masked_spanrD   rE   Úspec_aug_mask_idxÚdummy_mask_idxÚoffsetsrF   rG   s    `` `            @@r8   Ú_compute_mask_indicesrl   ‹   s­  ü€ ð0 #(Ñ€J�à�Q‚ÜÐAÓBÐBà�_Ò$ÜØ]Ð^iÐ]jØ& Ð&7°qð:ó
ð 	
ô �i‰i�n‰n˜QÓ×$Ñ$Ó&€G÷ð ð$ Ð%ð 	×ÑÓ×#Ñ# BÓ'×.Ñ.Ô0ä',¨ZÓ'8Ö9 !ŠoÒ9ð ô —H‘H˜j¨/Ð:Ä$ÔG€MØÐá1°/ÓBÐà˜aÒØÐà%ò 5ˆá1°,Ó?ˆô ŸI™I×,Ñ,Ü�I‰I�l k°A¡oÑ6Ó7¸ÐRWð -ó 
Ðô Ð Ó! QÒ&ð -¨qÑ0‰Nà.¨qÑ1ˆNäŸN™NØ¤§¡Ð(;¸oÑ(MÔUW×U]ÑU]Ô ^ÐaoÑ oÐpó
Ðð 	×!Ñ!Ð"3Õ4ð/5ô2 Ÿ™Ð"4Ó5Ðô Ÿ™Øš1ša ˜:Ñ&¨Ð5HÈ+Ð(VóÐð ,×3Ñ3°JÐ@SÐVaÑ@aÓbÐô �i‰i˜Ó$ T¨4² ]Ñ3€GÜ�o‰o˜g¨
Ð4GÈÐ'UÓV×^Ñ^ØÐ'¨+Ñ5ó€Gð ,¨gÑ5Ðð ×ÑÓ /°AÑ"5Ò5ØGVÐYZÑGZÐÐ-°À!Ñ0CÑCÑDô ×Ñ�mÐ%7¸¸BÔ?àÐùòw :s   Â$	I+Úfeatures_shapeÚnum_negativesÚmask_time_indicesc                 ód  — | \  }}t        j                  |«      }t        j                  |||ft         j                  ¬«      }|�|j	                  t
        «      nt        j                  | t
        ¬«      }t        |«      D ]­  }||   j                  «       dz
  }|||      }	t        j                  t        j                  |dz   «      dd…df   |dz   |f«      }
t         j                  j                  d||dz   |f¬«      }|||
k\  xx   dz  cc<   |	|   ||   ||   <   ||xx   ||z  z  cc<   Œ¯ |S )z>
    Sample `num_negatives` vectors from feature vectors.
    )r9   rK   NrJ   r    r   )Úsize)rN   rY   rV   r]   ÚastyperW   r\   rU   rS   r`   rO   Úrandint)rm   rn   ro   rc   rG   Úsequence_length_rangeÚsampled_negative_indicesÚ	batch_idxÚhighÚmapped_masked_indicesÚfeature_indicesÚsampled_indicess               r8   Ú_sample_negative_indicesr{     sT  € ð #1Ñ€J�ô ŸI™I oÓ6Ðô  "Ÿx™x¨z¸?ÈMÐ.ZÔbd×bjÑbjÔkÐð +<Ð*GÐ× Ñ ¤Ô&ÌRÏWÉWÐUcÔkoÔMpð ô ˜:Ó&ò Kˆ	Ø  Ñ+×/Ñ/Ó1°AÑ5ˆØ 5Ð6GÈ	Ñ6RÑ SÐäŸ/™/¬"¯)©)°D¸1±HÓ*=ºaÀ¸gÑ*FÈÐPQÉÐS`ÐHaÓbˆÜŸ)™)×+Ñ+¨A¨t¸4À!¹8À]Ð:SÐ+ÓTˆà˜¨?Ñ:Ó;¸qÑ@Ó;ð MbÐbqÑLrÐ  Ñ+Ð,=¸iÑ,HÑIð 	! Ó+¨y¸?Ñ/JÑJÔ+ðKð $Ð#r7   c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )ÚWav2Vec2NoLayerNormConvLayerc                 ód  •— t         ‰| �  «        |dkD  r|j                  |dz
     nd| _        |j                  |   | _        t        j                  | j                  | j                  |j                  |   |j                  |   |j                  ¬«      | _
        t        |j                     | _        y )Nr   r    ©Úkernel_sizeÚstrideÚbias)ÚsuperÚ__init__Úconv_dimÚin_conv_dimÚout_conv_dimr   ÚConv1dÚconv_kernelÚconv_strideÚ	conv_biasÚconvr
   Úfeat_extract_activationÚ
activation©ÚselfÚconfigÚlayer_idÚ	__class__s      €r8   r„   z%Wav2Vec2NoLayerNormConvLayer.__init__'  s—   ø€ Ü‰ÑÔØ<DÀqºL˜6Ÿ?™?¨8°a©<Ò8ÈaˆÔØ"ŸO™O¨HÑ5ˆÔä—I‘IØ×ÑØ×ÑØ×*Ñ*¨8Ñ4Ø×%Ñ% hÑ/Ø×!Ñ!ô
ˆŒ	ô ! ×!?Ñ!?Ñ@ˆ�r7   c                 óJ   — | j                  |«      }| j                  |«      }|S ©N)rŒ   rŽ   ©r�   r+   s     r8   Úforwardz$Wav2Vec2NoLayerNormConvLayer.forward5  s$   € ØŸ	™	 -Ó0ˆØŸ™¨Ó6ˆØÐr7   ©r   ©r/   r0   r1   r„   r—   Ú__classcell__©r“   s   @r8   r}   r}   &  s   ø„ õAör7   r}   c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )ÚWav2Vec2LayerNormConvLayerc                 ó°  •— t         ‰| �  «        |dkD  r|j                  |dz
     nd| _        |j                  |   | _        t        j                  | j                  | j                  |j                  |   |j                  |   |j                  ¬«      | _
        t        j                  | j                  d¬«      | _        t        |j                     | _        y )Nr   r    r   T)Úelementwise_affine)rƒ   r„   r…   r†   r‡   r   rˆ   r‰   rŠ   r‹   rŒ   Ú	LayerNormÚ
layer_normr
   r�   rŽ   r�   s      €r8   r„   z#Wav2Vec2LayerNormConvLayer.__init__<  s¯   ø€ Ü‰ÑÔØ<DÀqºL˜6Ÿ?™?¨8°a©<Ò8ÈaˆÔØ"ŸO™O¨HÑ5ˆÔä—I‘IØ×ÑØ×ÑØ×*Ñ*¨8Ñ4Ø×%Ñ% hÑ/Ø×!Ñ!ô
ˆŒ	ô Ÿ,™, t×'8Ñ'8ÈTÔRˆŒÜ  ×!?Ñ!?Ñ@ˆ�r7   c                 ó´   — | j                  |«      }|j                  dd«      }| j                  |«      }|j                  dd«      }| j                  |«      }|S )NéþÿÿÿrI   )rŒ   Ú	transposer¡   rŽ   r–   s     r8   r—   z"Wav2Vec2LayerNormConvLayer.forwardK  sV   € ØŸ	™	 -Ó0ˆà%×/Ñ/°°BÓ7ˆØŸ™¨Ó6ˆØ%×/Ñ/°°BÓ7ˆàŸ™¨Ó6ˆØÐr7   r˜   r™   r›   s   @r8   r�   r�   ;  s   ø„ õAör7   r�   c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )ÚWav2Vec2GroupNormConvLayerc                 óÆ  •— t         ‰| �  «        |dkD  r|j                  |dz
     nd| _        |j                  |   | _        t        j                  | j                  | j                  |j                  |   |j                  |   |j                  ¬«      | _
        t        |j                     | _        t        j                  | j                  | j                  d¬«      | _        y )Nr   r    r   T)Ú
num_groupsÚnum_channelsÚaffine)rƒ   r„   r…   r†   r‡   r   rˆ   r‰   rŠ   r‹   rŒ   r
   r�   rŽ   Ú	GroupNormr¡   r�   s      €r8   r„   z#Wav2Vec2GroupNormConvLayer.__init__W  s¹   ø€ Ü‰ÑÔØ<DÀqºL˜6Ÿ?™?¨8°a©<Ò8ÈaˆÔØ"ŸO™O¨HÑ5ˆÔä—I‘IØ×ÑØ×ÑØ×*Ñ*¨8Ñ4Ø×%Ñ% hÑ/Ø×!Ñ!ô
ˆŒ	ô ! ×!?Ñ!?Ñ@ˆŒäŸ,™,°$×2CÑ2CÐRV×RcÑRcÐlpÔqˆ�r7   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S r•   )rŒ   r¡   rŽ   r–   s     r8   r—   z"Wav2Vec2GroupNormConvLayer.forwardg  s2   € ØŸ	™	 -Ó0ˆØŸ™¨Ó6ˆØŸ™¨Ó6ˆØÐr7   r˜   r™   r›   s   @r8   r¦   r¦   V  s   ø„ õrö r7   r¦   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚWav2Vec2PositionalConvEmbeddingc                 ó¦  •— t         ‰| �  «        t        j                  |j                  |j                  |j
                  |j
                  dz  |j                  ¬«      | _        t        j                  j                  }t        t        j                  j                  d«      r$t        j                  j                  j                  }t        «       �r(dd l}|j                  j                  | j                  j                   d¬«      5   || j                  dd¬«      | _        d d d «       t        | j                  d«      rU| j                  j                  j                   j"                  }| j                  j                  j                   j$                  }n,| j                  j&                  }| j                  j(                  }|j                  j+                  | |«       |j                  j+                  | |«       n || j                  dd¬«      | _        t-        |j
                  «      | _        t0        |j2                     | _        y # 1 sw Y   �Œ'xY w)	Nr$   )r€   ÚpaddingÚgroupsÚweight_normr   )Úmodifier_rankÚweight)ÚnameÚdimÚparametrizations)rƒ   r„   r   rˆ   Úhidden_sizeÚnum_conv_pos_embeddingsÚnum_conv_pos_embedding_groupsrŒ   Úutilsr²   Úhasattrr·   r   Ú	deepspeedÚzeroÚGatheredParametersr´   Ú	original0Ú	original1Úweight_gÚweight_vÚregister_external_parameterÚWav2Vec2SamePadLayerr°   r
   r�   rŽ   )r�   r‘   r²   r½   rÂ   rÃ   r“   s         €r8   r„   z(Wav2Vec2PositionalConvEmbedding.__init__o  s¥  ø€ Ü‰ÑÔÜ—I‘IØ×ÑØ×ÑØ×6Ñ6Ø×2Ñ2°aÑ7Ø×7Ñ7ô
ˆŒ	ô —h‘h×*Ñ*ˆÜ”2—8‘8×,Ñ,¨mÔ<ÜŸ(™(×3Ñ3×?Ñ?ˆKä%Õ'Ûà—‘×2Ñ2°4·9±9×3CÑ3CÐSTÐ2ÓUñ IÙ'¨¯	©	¸ÀaÔH�”	÷Iä�t—y‘yÐ"4Ô5ØŸ9™9×5Ñ5×<Ñ<×FÑF�ØŸ9™9×5Ñ5×<Ñ<×FÑF‘àŸ9™9×-Ñ-�ØŸ9™9×-Ñ-�Ø�N‰N×6Ñ6°t¸XÔFØ�N‰N×6Ñ6°t¸XÕFá# D§I¡I°HÀ!ÔDˆDŒIä+¨F×,JÑ,JÓKˆŒÜ  ×!?Ñ!?Ñ@ˆ�÷Iñ Iús   ÄIÉIc                 ó´   — |j                  dd«      }| j                  |«      }| j                  |«      }| j                  |«      }|j                  dd«      }|S ©Nr    r$   )r¤   rŒ   r°   rŽ   r–   s     r8   r—   z'Wav2Vec2PositionalConvEmbedding.forward�  sV   € Ø%×/Ñ/°°1Ó5ˆàŸ	™	 -Ó0ˆØŸ™ ]Ó3ˆØŸ™¨Ó6ˆà%×/Ñ/°°1Ó5ˆØÐr7   r™   r›   s   @r8   r®   r®   n  s   ø„ ôAöBr7   r®   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )rÅ   c                 óP   •— t         ‰| �  «        |dz  dk(  rd| _        y d| _        y )Nr$   r   r    )rƒ   r„   Únum_pad_remove)r�   r¹   r“   s     €r8   r„   zWav2Vec2SamePadLayer.__init__œ  s)   ø€ Ü‰ÑÔØ#:¸QÑ#>À!Ò#C˜aˆÕÈˆÕr7   c                 óV   — | j                   dkD  r|d d …d d …d | j                    …f   }|S ©Nr   )rÊ   r–   s     r8   r—   zWav2Vec2SamePadLayer.forward   s6   € Ø×Ñ Ò"Ø)ª!ªQÐ0F°4×3FÑ3FÐ2FÐ0FÐ*FÑGˆMØÐr7   r™   r›   s   @r8   rÅ   rÅ   ›  s   ø„ ôKör7   rÅ   c                   ó.   ‡ — e Zd ZdZˆ fd„Zd„ Zd„ Zˆ xZS )ÚWav2Vec2FeatureEncoderz.Construct the features from raw audio waveformc           	      óØ  •— t         ‰| �  «        |j                  dk(  rDt        |d¬«      gt	        |j
                  dz
  «      D �cg c]  }t        ||dz   ¬«      ‘Œ c}z   }nV|j                  dk(  r.t	        |j
                  «      D �cg c]  }t        ||¬«      ‘Œ }}nt        d|j                  › d�«      ‚t        j                  |«      | _        d| _        d	| _        y c c}w c c}w )
NÚgroupr   )r’   r    Úlayerz`config.feat_extract_norm` is z), but has to be one of ['group', 'layer']FT)rƒ   r„   Úfeat_extract_normr¦   rU   Únum_feat_extract_layersr}   r�   rM   r   Ú
ModuleListÚconv_layersÚgradient_checkpointingÚ_requires_grad)r�   r‘   ÚirÕ   r“   s       €r8   r„   zWav2Vec2FeatureEncoder.__init__©  sñ   ø€ Ü‰ÑÔà×#Ñ# wÒ.Ü5°fÀqÔIÐJÜNSÐTZ×TrÑTrÐuvÑTvÓNwöNØIJÔ,¨V¸aÀ!¹eÖDòNñ ‰Kð ×%Ñ%¨Ò0äHMÈf×NlÑNlÓHmöØCDÔ*¨6¸AÖ>ðˆKñ ô Ø0°×1IÑ1IÐ0JÐJsÐtóð ô Ÿ=™=¨Ó5ˆÔØ&+ˆÔ#Ø"ˆÕùòNùòs   ÁC"Â	C'c                 óJ   — | j                  «       D ]	  }d|_        Œ d| _        y ©NF)Ú
parametersÚrequires_gradr×   ©r�   Úparams     r8   Ú_freeze_parametersz)Wav2Vec2FeatureEncoder._freeze_parameters¼  s(   € Ø—_‘_Ó&ò 	(ˆEØ"'ˆEÕð	(à#ˆÕr7   c                 ó
  — |d d …d f   }| j                   r| j                  rd|_        | j                  D ]K  }| j                   r5| j                  r)| j                  r| j                  |j                  |«      }ŒD ||«      }ŒM |S )NT)r×   ÚtrainingrÜ   rÕ   rÖ   Ú_gradient_checkpointing_funcÚ__call__)r�   Úinput_valuesr+   Ú
conv_layers       r8   r—   zWav2Vec2FeatureEncoder.forwardÁ  s…   € Ø$¢Q¨ WÑ-ˆð ×Ò 4§=¢=Ø*.ˆMÔ'à×*Ñ*ò 	:ˆJØ×"Ò" t×'BÒ'BÀtÇ}Â}Ø $× AÑ AØ×'Ñ'Ø!ó!‘ñ
 !+¨=Ó 9‘ð	:ð Ðr7   )r/   r0   r1   r2   r„   rß   r—   rš   r›   s   @r8   rÎ   rÎ   ¦  s   ø„ Ù8ô#ò&$ö
r7   rÎ   c                   ó   ‡ — e Zd Zˆ fd„Zˆ xZS )ÚWav2Vec2FeatureExtractorc                 óÐ   •— t         ‰| �  |«       t        j                  d| j                  j
                  › d| j                  j                  d   j
                  › d�t        «       y )NzThe class `zD` has been depreciated and will be removed in Transformers v5. Use `r   z
` instead.)rƒ   r„   ÚwarningsÚwarnr“   r/   Ú	__bases__ÚFutureWarning©r�   r‘   r“   s     €r8   r„   z!Wav2Vec2FeatureExtractor.__init__Õ  s[   ø€ Ü‰Ñ˜Ô Ü�‰Ø˜$Ÿ.™.×1Ñ1Ð2ð 3à—N‘N×,Ñ,¨QÑ/×8Ñ8Ð9¸ðEô õ		
r7   )r/   r0   r1   r„   rš   r›   s   @r8   rç   rç   Ô  s   ø„ ÷
ð 
r7   rç   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚWav2Vec2FeatureProjectionc                 ó4  •— t         ‰| �  «        t        j                  |j                  d   |j
                  ¬«      | _        t        j                  |j                  d   |j                  «      | _	        t        j                  |j                  «      | _        y )NrI   ©Úeps)rƒ   r„   r   r    r…   Úlayer_norm_epsr¡   ÚLinearr¸   Ú
projectionÚDropoutÚfeat_proj_dropoutÚdropoutrí   s     €r8   r„   z"Wav2Vec2FeatureProjection.__init__à  sf   ø€ Ü‰ÑÔÜŸ,™, v§¡°rÑ':À×@UÑ@UÔVˆŒÜŸ)™) F§O¡O°BÑ$7¸×9KÑ9KÓLˆŒÜ—z‘z &×":Ñ":Ó;ˆ�r7   c                 óp   — | j                  |«      }| j                  |«      }| j                  |«      }||fS r•   )r¡   rõ   rø   )r�   r+   Únorm_hidden_statess      r8   r—   z!Wav2Vec2FeatureProjection.forwardæ  s:   € à!Ÿ_™_¨]Ó;ÐØŸ™Ð(:Ó;ˆØŸ™ ]Ó3ˆØÐ0Ð0Ð0r7   r™   r›   s   @r8   rï   rï   ß  s   ø„ ô<ö1r7   rï   c                   ó†  ‡ — e Zd ZdZ	 	 	 	 	 ddededededededee   fˆ fd	„Z	d
e
j                  dedefd„Z	 	 	 	 	 dde
j                  dee
j                     deee
j                        dee
j                     dee
j                     dedee
j                  ee
j                     eee
j                        f   fd„Zˆ xZS )ÚWav2Vec2Attentionz=Multi-headed attention from 'Attention Is All You Need' paperÚ	embed_dimÚ	num_headsrø   Ú
is_decoderr‚   Ú	is_causalr‘   c                 ó
  •— t         ‰| �  «        || _        || _        || _        ||z  | _        || _        | j
                  |z  | j                  k7  rt        d| j                  › d|› d�«      ‚| j
                  dz  | _        || _	        || _
        t        j                  |||¬«      | _        t        j                  |||¬«      | _        t        j                  |||¬«      | _        t        j                  |||¬«      | _        y )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: z).g      à¿)r‚   )rƒ   r„   rý   rþ   rø   Úhead_dimr‘   rM   Úscalingrÿ   r   r   rô   Úk_projÚv_projÚq_projÚout_proj)	r�   rý   rþ   rø   rÿ   r‚   r   r‘   r“   s	           €r8   r„   zWav2Vec2Attention.__init__ò  sä   ø€ ô 	‰ÑÔØ"ˆŒØ"ˆŒØˆŒØ! YÑ.ˆŒØˆŒà�M‰M˜IÑ%¨$¯.©.Ò8ÜØMÈdÏnÉnÐM]Ø$ Y K¨rð3óð ð —}‘} dÑ*ˆŒØ$ˆŒØ"ˆŒä—i‘i 	¨9¸4Ô@ˆŒÜ—i‘i 	¨9¸4Ô@ˆŒÜ—i‘i 	¨9¸4Ô@ˆŒÜŸ	™	 )¨Y¸TÔBˆ�r7   ÚtensorÚseq_lenÚbszc                 óŽ   — |j                  ||| j                  | j                  «      j                  dd«      j	                  «       S rÇ   )Úviewrþ   r  r¤   Ú
contiguous©r�   r  r	  r
  s       r8   Ú_shapezWav2Vec2Attention._shape  s7   € Ø�{‰{˜3 ¨¯©¸¿¹ÓG×QÑQÐRSÐUVÓW×bÑbÓdÐdr7   r+   Úkey_value_statesÚpast_key_valuer<   Úlayer_head_maskÚoutput_attentionsr>   c                 ó
  — |du}|j                  «       \  }}	}
| j                  |«      | j                  z  }|r0|�.|d   j                  d   |j                  d   k(  r|d   }|d   }�n
|rE| j	                  | j                  |«      d|«      }| j	                  | j                  |«      d|«      }nÃ|�}| j	                  | j                  |«      d|«      }| j	                  | j                  |«      d|«      }t        j                  |d   |gd¬«      }t        j                  |d   |gd¬«      }nD| j	                  | j                  |«      d|«      }| j	                  | j                  |«      d|«      }| j                  r||f}|| j                  z  d| j                  f} | j	                  ||	|«      j                  |Ž } |j                  |Ž } |j                  |Ž }|j                  d«      }t        j                  ||j                  dd«      «      }|j                  «       || j                  z  |	|fk7  r/t!        d|| j                  z  |	|f› d|j                  «       › �«      ‚|�{|j                  «       |d|	|fk7  r#t!        d	|d|	|f› d|j                  «       › �«      ‚|j                  || j                  |	|«      |z   }|j                  || j                  z  |	|«      }t"        j$                  j'                  |d¬«      }|�›|j                  «       | j                  fk7  r*t!        d
| j                  f› d|j                  «       › �«      ‚|j                  dddd«      |j                  || j                  |	|«      z  }|j                  || j                  z  |	|«      }|r?|j                  || j                  |	|«      }|j                  || j                  z  |	|«      }nd}t"        j$                  j)                  || j(                  | j*                  ¬«      }t        j                  ||«      }|j                  «       || j                  z  |	| j                  fk7  r9t!        d|| j                  z  |	| j                  f› d|j                  «       › �«      ‚|j                  || j                  |	| j                  «      }|j                  dd«      }|j                  ||	| j,                  «      }| j/                  |«      }|||fS )ú#Input shape: Batch x Time x ChannelNr   r$   r    rI   ©r¶   z$Attention weights should be of size ú	, but is z!Attention mask should be of size z/Head mask for a single layer should be of size )Úprá   ú `attn_output` should be of size )rq   r  r  r9   r  r  r  r3   Úcatrÿ   rþ   r  r  ra   Úbmmr¤   rM   r   Ú
functionalÚsoftmaxrø   rá   rý   r  )r�   r+   r  r  r<   r  r  Úis_cross_attentionr
  Útgt_lenrd   Úquery_statesÚ
key_statesÚvalue_statesÚ
proj_shapeÚsrc_lenÚattn_weightsÚattn_weights_reshapedÚ
attn_probsÚattn_outputs                       r8   r—   zWav2Vec2Attention.forward  s  € ð .°TÐ9Ðà'×,Ñ,Ó.‰ˆˆW�að —{‘{ =Ó1°D·L±LÑ@ˆñ ØÐ*Ø˜qÑ!×'Ñ'¨Ñ*Ð.>×.DÑ.DÀQÑ.GÒGð (¨Ñ*ˆJØ)¨!Ñ,ŠLÙàŸ™ T§[¡[Ð1AÓ%BÀBÈÓLˆJØŸ;™; t§{¡{Ð3CÓ'DÀbÈ#ÓN‰LØÐ'àŸ™ T§[¡[°Ó%?ÀÀSÓIˆJØŸ;™; t§{¡{°=Ó'AÀ2ÀsÓKˆLÜŸ™ N°1Ñ$5°zÐ#BÈÔJˆJÜ Ÿ9™9 n°QÑ&7¸Ð%FÈAÔN‰Lð Ÿ™ T§[¡[°Ó%?ÀÀSÓIˆJØŸ;™; t§{¡{°=Ó'AÀ2ÀsÓKˆLà�?Š?ð )¨,Ð7ˆNà˜DŸN™NÑ*¨B°·±Ð>ˆ
ØC�t—{‘{ <°¸#Ó>×CÑCÀZÐPˆØ'�Z×'Ñ'¨Ð4ˆ
Ø+�|×+Ñ+¨ZÐ8ˆà—/‘/ !Ó$ˆÜ—y‘y ¨z×/CÑ/CÀAÀqÓ/IÓJˆà×ÑÓ 3¨¯©Ñ#7¸À'Ð"JÒJÜØ6¸¸d¿n¹nÑ8LÈgÐW^Ð7_Ð6`ð aØ ×%Ñ%Ó'Ð(ð*óð ð
 Ð%Ø×"Ñ"Ó$¨¨a°¸'Ð(BÒBÜ Ø7¸¸aÀÈ'Ð8RÐ7SÐS\Ð]k×]pÑ]pÓ]rÐ\sÐtóð ð (×,Ñ,¨S°$·.±.À'È7ÓSÐVdÑdˆLØ'×,Ñ,¨S°4·>±>Ñ-AÀ7ÈGÓTˆLä—}‘}×,Ñ,¨\¸rÐ,ÓBˆàÐ&Ø×#Ñ#Ó%¨$¯.©.Ð):Ò:Ü ØEÀtÇ~Á~ÐFWÐEXð YØ'×,Ñ,Ó.Ð/ð1óð ð +×/Ñ/°°2°q¸!Ó<¸|×?PÑ?PÐQTÐVZ×VdÑVdÐfmÐovÓ?wÑwˆLØ'×,Ñ,¨S°4·>±>Ñ-AÀ7ÈGÓTˆLáð
 %1×$5Ñ$5°c¸4¿>¹>È7ÐT[Ó$\Ð!Ø0×5Ñ5°c¸D¿N¹NÑ6JÈGÐU\Ó]‰Là$(Ð!ä—]‘]×*Ñ*¨<¸4¿<¹<ÐRV×R_ÑR_Ð*Ó`ˆ
ä—i‘i 
¨LÓ9ˆà×ÑÓ #¨¯©Ñ"6¸ÀÇÁÐ!OÒOÜØ2°C¸$¿.¹.Ñ4HÈ'ÐSW×S`ÑS`Ð3aÐ2bð cØ×$Ñ$Ó&Ð'ð)óð ð
 "×&Ñ& s¨D¯N©N¸GÀTÇ]Á]ÓSˆØ!×+Ñ+¨A¨qÓ1ˆð "×)Ñ)¨#¨w¸¿¹ÓGˆà—m‘m KÓ0ˆàÐ1°>ÐAÐAr7   )ç        FTFN©NNNNF)r/   r0   r1   r2   rB   ÚfloatrW   r   r!   r„   r3   ÚTensorr  r   r—   rš   r›   s   @r8   rü   rü   ï  sM  ø„ ÙGð Ø ØØØ+/ñCàðCð ðCð ð	Cð
 ðCð ðCð ðCð ˜Ñ(õCð>e˜UŸ\™\ð e°Cð e¸có eð 48Ø8<Ø15Ø26Ø"'ñvBà—|‘|ðvBð # 5§<¡<Ñ0ðvBð !  u§|¡|Ñ!4Ñ5ð	vBð
 ! §¡Ñ.ðvBð " %§,¡,Ñ/ðvBð  ðvBð 
ˆu�|‰|˜X e§l¡lÑ3°X¸eÀEÇLÁLÑ>QÑ5RÐRÑ	S÷vBr7   rü   c                   óV  ‡ — e Zd ZdZˆ fd„Zdej                  dedefd„Z	 	 	 	 	 ddej                  de	ej                     d	e	e
ej                        d
e	ej                     de	ej                     dede
ej                  e	ej                     e	e
ej                        f   fd„Zˆ xZS )ÚWav2Vec2FlashAttention2aL  
    Wav2Vec2 flash attention module. This module inherits from `Wav2Vec2Attention` as the weights of the module stays
    untouched. The only required change would be on the forward pass where it needs to correctly call the public API of
    flash attention and deal with padding tokens in case the input contains any of them.
    c                 óB   •— t        ‰| �  |i |¤Ž t        «       | _        y r•   )rƒ   r„   r   Ú_flash_attn_uses_top_left_mask)r�   ÚargsÚkwargsr“   s      €r8   r„   z Wav2Vec2FlashAttention2.__init__•  s#   ø€ Ü‰Ñ˜$Ð) &Ò)ô
 /PÓ.QˆÕ+r7   r  r	  r
  c                 óR   — |j                  ||| j                  | j                  «      S r•   )r  rþ   r  r  s       r8   Ú_reshapez Wav2Vec2FlashAttention2._reshape�  s   € Ø�{‰{˜3 ¨¯©¸¿¹ÓGÐGr7   r+   r  r  r<   r  r  r>   c           
      óÎ  — |rt        d«      ‚|d u}|j                  «       \  }}	}
| j                  | j                  |«      d|«      }|rP|�N|d   j                  d   |j                  d   k(  r,|d   j                  dd«      }|d   j                  dd«      }�n*|rE| j                  | j                  |«      d|«      }| j                  | j                  |«      d|«      }nã|��| j                  | j                  |«      d|«      }| j                  | j                  |«      d|«      }t        j                  |d   j                  dd«      |gd¬«      }t        j                  |d   j                  dd«      |gd¬«      }nD| j                  | j                  |«      d|«      }| j                  | j                  |«      d|«      }| j                  r$|j                  dd«      |j                  dd«      f}|j                  d   }|�||d   j                  d   z  }|j                  }|t        j                  k(  rÂt        j                  «       rt        j                  «       }nMt        | j                   d«      r| j                   j"                  }n | j                  j$                  j                  }t&        j)                  d	|› d
�«       |j+                  |«      }|j+                  |«      }|j+                  |«      }t-        |||||	| j.                  r| j0                  nd| j2                  | j4                  ¬«      }|j7                  ||	d«      }| j9                  |«      }|sd }||fS )NzDWav2Vec2FlashAttention2 attention does not support output_attentionsrI   r   r$   r    r  r£   Ú_pre_quantization_dtypez¾The input hidden states seems to be silently casted in float32, this might be related to the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in ú.r)  )rø   r   Úuse_top_left_mask)rM   rq   r4  r  r9   r¤   r  r  r3   r  rÿ   rK   Úfloat32Úis_autocast_enabledÚget_autocast_gpu_dtyper¼   r‘   r6  r´   ÚloggerÚwarning_onceÚtor#   rá   rø   r   r0  ra   r  )r�   r+   r  r  r<   r  r  r  r
  Úq_lenrd   r   r!  r"  Ú
kv_seq_lenÚinput_dtypeÚtarget_dtyper(  r%  s                      r8   r—   zWav2Vec2FlashAttention2.forward   s0  € ñ ÜÐcÓdÐdð .°TÐ9Ðà%×*Ñ*Ó,‰ˆˆU�Að —}‘} T§[¡[°Ó%?ÀÀSÓIˆñ ØÐ*Ø˜qÑ!×'Ñ'¨Ñ*Ð.>×.DÑ.DÀQÑ.GÒGð (¨Ñ*×4Ñ4°Q¸Ó:ˆJØ)¨!Ñ,×6Ñ6°q¸!Ó<ŠLÙàŸ™ t§{¡{Ð3CÓ'DÀbÈ#ÓNˆJØŸ=™=¨¯©Ð5EÓ)FÈÈCÓP‰LØÐ'àŸ™ t§{¡{°=Ó'AÀ2ÀsÓKˆJØŸ=™=¨¯©°]Ó)CÀRÈÓMˆLÜŸ™ N°1Ñ$5×$?Ñ$?ÀÀ1Ó$EÀzÐ#RÐXYÔZˆJÜ Ÿ9™9 n°QÑ&7×&AÑ&AÀ!ÀQÓ&GÈÐ%VÐ\]Ô^‰Lð Ÿ™ t§{¡{°=Ó'AÀ2ÀsÓKˆJØŸ=™=¨¯©°]Ó)CÀRÈÓMˆLà�?Š?ð )×2Ñ2°1°aÓ8¸,×:PÑ:PÐQRÐTUÓ:VÐWˆNà×%Ñ% bÑ)ˆ
ØÐ%Ø˜.¨Ñ+×1Ñ1°"Ñ5Ñ5ˆJð #×(Ñ(ˆØœ%Ÿ-™-Ò'Ü×(Ñ(Ô*Ü$×;Ñ;Ó=‘ä˜Ÿ™Ð&?Ô@Ø#Ÿ{™{×BÑB‘à#Ÿ{™{×1Ñ1×7Ñ7�ä×Ñðà �> ð$ôð (Ÿ?™?¨<Ó8ˆLØ#Ÿ™ |Ó4ˆJØ'Ÿ?™?¨<Ó8ˆLä.ØØØØØØ$(§M¢M�D—L’L°sØ—n‘nØ"×AÑAô	
ˆð "×)Ñ)¨#¨u°bÓ9ˆØ—m‘m KÓ0ˆá ØˆLà˜L¨.Ð8Ð8r7   r*  )r/   r0   r1   r2   r„   r3   r,  rB   r4  r   r   rW   r—   rš   r›   s   @r8   r.  r.  Ž  sæ   ø„ ñôRðH˜uŸ|™|ð H°cð HÀó Hð 48Ø8<Ø15Ø26Ø"'ñi9à—|‘|ði9ð # 5§<¡<Ñ0ði9ð !  u§|¡|Ñ!4Ñ5ð	i9ð
 ! §¡Ñ.ði9ð " %§,¡,Ñ/ði9ð  ði9ð 
ˆu�|‰|˜X e§l¡lÑ3°X¸eÀEÇLÁLÑ>QÑ5RÐRÑ	S÷i9r7   r.  c                   ó$  ‡ — e Zd Z	 	 	 	 	 d	dej                  deej                     deeej                        deej                     deej                     dedeej                  eej                     eeej                        f   fˆ fd„Zˆ xZ	S )
ÚWav2Vec2SdpaAttentionr+   r  r  r<   r  r  r>   c                 óz  •— |s|�*t         j                  d«       t        ‰| �  ||||||¬«      S |du}|j	                  «       \  }}	}
| j                  |«      }|r0|�.|d   j                  d   |j                  d   k(  r|d   }|d   }�n
|rE| j                  | j                  |«      d|«      }| j                  | j                  |«      d|«      }nÃ|�}| j                  | j                  |«      d|«      }| j                  | j                  |«      d|«      }t        j                  |d   |gd¬«      }t        j                  |d   |gd¬«      }nD| j                  | j                  |«      d|«      }| j                  | j                  |«      d|«      }| j                  r||f}| j                  ||	|«      }| j                  r	|€|	dkD  rd	nd
}t        j                  j                  j!                  ||||| j"                  r| j$                  nd|¬«      }|j	                  «       || j&                  |	| j(                  fk7  r7t+        d|| j&                  |	| j(                  f› d|j	                  «       › �«      ‚|j-                  dd«      }|j/                  ||	| j0                  «      }| j3                  |«      }|d|fS )r  Na«  Wav2Vec2Model is using Wav2Vec2SdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True` or `layer_head_mask` not None. Falling back to the manual attention implementation, but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.)r  r  r<   r  r  r   r$   r    rI   r  TFr)  )Ú	attn_maskÚ	dropout_pr   r  r  )r<  r=  rƒ   r—   rq   r  r9   r  r  r  r3   r  rÿ   r   r   r  Úscaled_dot_product_attentionrá   rø   rþ   r  rM   r¤   ra   rý   r  )r�   r+   r  r  r<   r  r  r  r
  r  rd   r   r!  r"  r   r(  r“   s                   €r8   r—   zWav2Vec2SdpaAttention.forward  sÕ  ø€ ñ  Ð ;ä×Ñðlôô ‘7‘?ØØ!1Ø-Ø-Ø /Ø"3ð #ó ð ð .°TÐ9Ðà'×,Ñ,Ó.‰ˆˆW�að —{‘{ =Ó1ˆñ ØÐ*Ø˜qÑ!×'Ñ'¨Ñ*Ð.>×.DÑ.DÀQÑ.GÒGð (¨Ñ*ˆJØ)¨!Ñ,ŠLÙàŸ™ T§[¡[Ð1AÓ%BÀBÈÓLˆJØŸ;™; t§{¡{Ð3CÓ'DÀbÈ#ÓN‰LØÐ'àŸ™ T§[¡[°Ó%?ÀÀSÓIˆJØŸ;™; t§{¡{°=Ó'AÀ2ÀsÓKˆLÜŸ™ N°1Ñ$5°zÐ#BÈÔJˆJÜ Ÿ9™9 n°QÑ&7¸Ð%FÈAÔN‰Lð Ÿ™ T§[¡[°Ó%?ÀÀSÓIˆJØŸ;™; t§{¡{°=Ó'AÀ2ÀsÓKˆLà�?Š?ð )¨,Ð7ˆNà—{‘{ <°¸#Ó>ˆð
 !ŸNšN¨~Ð/EÈ'ÐTUÊ+‘DÐ[`ˆ	ô —h‘h×)Ñ)×FÑFØØØØ$Ø&*§m¢m�d—l’l¸Øð Gó 
ˆð ×ÑÓ # t§~¡~°wÀÇÁÐ!NÒNÜØ2°C¸¿¹ÈÐRV×R_ÑR_Ð3`Ð2að bØ×$Ñ$Ó&Ð'ð)óð ð
 "×+Ñ+¨A¨qÓ1ˆð "×)Ñ)¨#¨w¸¿¹ÓGˆà—m‘m KÓ0ˆà˜D .Ð0Ð0r7   r*  )
r/   r0   r1   r3   r,  r   r   rW   r—   rš   r›   s   @r8   rD  rD    s¿   ø„ ð
 48Ø8<Ø15Ø26Ø"'ñf1à—|‘|ðf1ð # 5§<¡<Ñ0ðf1ð !  u§|¡|Ñ!4Ñ5ð	f1ð
 ! §¡Ñ.ðf1ð " %§,¡,Ñ/ðf1ð  ðf1ð 
ˆu�|‰|˜X e§l¡lÑ3°X¸eÀEÇLÁLÑ>QÑ5RÐRÑ	S÷f1ñ f1r7   rD  )ÚeagerÚsdpaÚflash_attention_2c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚWav2Vec2FeedForwardc                 óö  •— t         ‰| �  «        t        j                  |j                  «      | _        t        j                  |j                  |j                  «      | _	        t        |j                  t        «      rt        |j                     | _        n|j                  | _        t        j                  |j                  |j                  «      | _        t        j                  |j                   «      | _        y r•   )rƒ   r„   r   rö   Úactivation_dropoutÚintermediate_dropoutrô   r¸   Úintermediate_sizeÚintermediate_denseÚ
isinstanceÚ
hidden_actÚstrr
   Úintermediate_act_fnÚoutput_denseÚhidden_dropoutÚoutput_dropoutrí   s     €r8   r„   zWav2Vec2FeedForward.__init__  s«   ø€ Ü‰ÑÔÜ$&§J¡J¨v×/HÑ/HÓ$IˆÔ!ä"$§)¡)¨F×,>Ñ,>À×@XÑ@XÓ"YˆÔÜ�f×'Ñ'¬Ô-Ü'-¨f×.?Ñ.?Ñ'@ˆDÕ$à'-×'8Ñ'8ˆDÔ$äŸI™I f×&>Ñ&>À×@RÑ@RÓSˆÔÜ Ÿj™j¨×)>Ñ)>Ó?ˆÕr7   c                 ó°   — | j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j	                  |«      }|S r•   )rR  rV  rP  rW  rY  r–   s     r8   r—   zWav2Vec2FeedForward.forwardŒ  sX   € Ø×/Ñ/°Ó>ˆØ×0Ñ0°Ó?ˆØ×1Ñ1°-Ó@ˆà×)Ñ)¨-Ó8ˆØ×+Ñ+¨MÓ:ˆØÐr7   r™   r›   s   @r8   rM  rM  ~  s   ø„ ô@ör7   rM  c                   ó&   ‡ — e Zd Zˆ fd„Zdd„Zˆ xZS )ÚWav2Vec2EncoderLayerc                 óÈ  •— t         ‰| �  «        t        |j                     |j                  |j
                  |j                  d¬«      | _        t        j                  |j                  «      | _        t        j                  |j                  |j                  ¬«      | _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _        y )NF©rý   rþ   rø   rÿ   rñ   )rƒ   r„   ÚWAV2VEC2_ATTENTION_CLASSESÚ_attn_implementationr¸   Únum_attention_headsÚattention_dropoutÚ	attentionr   rö   rX  rø   r    ró   r¡   rM  Úfeed_forwardÚfinal_layer_normrí   s     €r8   r„   zWav2Vec2EncoderLayer.__init__—  s¥   ø€ Ü‰ÑÔÜ3°F×4OÑ4OÑPØ×(Ñ(Ø×0Ñ0Ø×,Ñ,Øô	
ˆŒô —z‘z &×"7Ñ"7Ó8ˆŒÜŸ,™, v×'9Ñ'9¸v×?TÑ?TÔUˆŒÜ/°Ó7ˆÔÜ "§¡¨V×-?Ñ-?ÀV×EZÑEZÔ [ˆÕr7   c                 óè   — |}| j                  |||¬«      \  }}}| j                  |«      }||z   }| j                  |«      }|| j                  |«      z   }| j	                  |«      }|f}|r||fz  }|S ©N©r<   r  )rc  rø   r¡   rd  re  ©r�   r+   r<   r  Úattn_residualr%  rd   Úoutputss           r8   r—   zWav2Vec2EncoderLayer.forward¥  s’   € Ø%ˆØ)-¯©Ø¨.ÐL]ð *8ó *
Ñ&ˆ�| Qð Ÿ™ ]Ó3ˆØ%¨Ñ5ˆàŸ™¨Ó6ˆØ%¨×(9Ñ(9¸-Ó(HÑHˆØ×-Ñ-¨mÓ<ˆà Ð"ˆáØ˜�Ñ&ˆGàˆr7   rÚ   r™   r›   s   @r8   r\  r\  –  s   ø„ ô\÷r7   r\  c                   óf   ‡ — e Zd Zˆ fd„Z	 	 ddej
                  deej
                     defd„Zˆ xZ	S )Ú#Wav2Vec2EncoderLayerStableLayerNormc                 ó  •— t         ‰| �  «        t        |j                     |j                  |j
                  |j                  d¬«      | _        t        j                  |j                  «      | _        t        j                  |j                  |j                  ¬«      | _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _        t%        |dd «      �t'        |«      | _        y d | _        y )NFr^  rñ   Úadapter_attn_dim)rƒ   r„   r_  r`  r¸   ra  rb  rc  r   rö   rX  rø   r    ró   r¡   rM  rd  re  ÚgetattrÚWav2Vec2AttnAdapterLayerÚadapter_layerrí   s     €r8   r„   z,Wav2Vec2EncoderLayerStableLayerNorm.__init__º  sÊ   ø€ Ü‰ÑÔÜ3°F×4OÑ4OÑPØ×(Ñ(Ø×0Ñ0Ø×,Ñ,Øô	
ˆŒô —z‘z &×"7Ñ"7Ó8ˆŒÜŸ,™, v×'9Ñ'9¸v×?TÑ?TÔUˆŒÜ/°Ó7ˆÔÜ "§¡¨V×-?Ñ-?ÀV×EZÑEZÔ [ˆÔä�6Ð-¨tÓ4Ð@Ü!9¸&Ó!AˆDÕà!%ˆDÕr7   r+   r<   r  c                 ó$  — |}| j                  |«      }| j                  |||¬«      \  }}}| j                  |«      }||z   }|| j                  | j	                  |«      «      z   }| j
                  �|| j                  |«      z   }|f}|r||fz  }|S rg  )r¡   rc  rø   rd  re  rr  ri  s           r8   r—   z+Wav2Vec2EncoderLayerStableLayerNorm.forwardÌ  s±   € ð &ˆØŸ™¨Ó6ˆØ)-¯©Ø¨.ÐL]ð *8ó *
Ñ&ˆ�| Qð Ÿ™ ]Ó3ˆØ%¨Ñ5ˆØ%¨×(9Ñ(9¸$×:OÑ:OÐP]Ó:^Ó(_Ñ_ˆà×ÑÐ)Ø)¨D×,>Ñ,>¸}Ó,MÑMˆMà Ð"ˆáØ˜�Ñ&ˆGàˆr7   rÚ   )
r/   r0   r1   r„   r3   r,  r   rW   r—   rš   r›   s   @r8   rm  rm  ¹  s>   ø„ ô&ð* 26Ø"'ñ	à—|‘|ðð ! §¡Ñ.ðð  ÷	r7   rm  c                   ór   ‡ — e Zd Zˆ fd„Z	 	 	 	 ddej
                  deej                     dededef
d„Z	ˆ xZ
S )	ÚWav2Vec2Encoderc                 óÀ  •— t         ‰| �  «        || _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _	        t        j                  |j                  «      | _        t        j                  t        |j                  «      D �cg c]  }t!        |«      ‘Œ c}«      | _        d| _        |j&                  dk(  | _        y c c}w ©Nrñ   FrK  )rƒ   r„   r‘   r®   Úpos_conv_embedr   r    r¸   ró   r¡   rö   rX  rø   rÔ   rU   Únum_hidden_layersr\  ÚlayersrÖ   r`  Ú_use_flash_attention_2©r�   r‘   rd   r“   s      €r8   r„   zWav2Vec2Encoder.__init__ç  s¥   ø€ Ü‰ÑÔØˆŒÜ=¸fÓEˆÔÜŸ,™, v×'9Ñ'9¸v×?TÑ?TÔUˆŒÜ—z‘z &×"7Ñ"7Ó8ˆŒÜ—m‘mÌ5ÐQW×QiÑQiÓKjÖ$kÀaÔ%9¸&Õ%AÒ$kÓlˆŒØ&+ˆÔ#Ø&,×&AÑ&AÐEXÑ&XˆÕ#ùò %ló   Â!Cr+   r<   r  Úoutput_hidden_statesÚreturn_dictc                 ó4  — |rdnd }|rdnd }|�Ý|j                  d«      j                  dd|j                  d   «      }d|| <   | j                  r|�d|v r|nd }n‘d|d d …d d d d …f   j	                  |j
                  ¬«      z
  }|t        j                  |j
                  «      j                  z  }|j                  |j                  d   d|j                  d   |j                  d   «      }| j                  |«      }	||	z   }| j                  |«      }| j                  |«      }t        «       xs t        | «      }
| j                  D ]£  }|r||fz   }t        j                   g «      }| j"                  r|| j$                  j&                  k  rdnd	}|r|
rG| j(                  r+| j"                  r| j+                  |j,                  |||«      }n ||||¬
«      }|d   }|rd}|sŒ›|d   fz   }Œ¥ |r||fz   }|st/        d„ |||fD «       «      S t1        |||¬«      S )Nr6   rI   r    r$   r   ç      ð?rJ   TFrh  ©NNc              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wr•   r6   ©Ú.0Úvs     r8   ú	<genexpr>z*Wav2Vec2Encoder.forward.<locals>.<genexpr>3  ó   è ø€ Òm˜qÐ_`Ñ_lœÑmùó   ‚Š©Úlast_hidden_stater+   r,   )Ú	unsqueezeÚrepeatr9   r{  r>  rK   r3   ÚfinfoÚminÚexpandrx  r¡   rø   r   r   rz  rP   rá   r‘   Ú	layerdroprÖ   râ   rã   Útupler   ©r�   r+   r<   r  r~  r  Úall_hidden_statesÚall_self_attentionsÚexpand_attention_maskÚposition_embeddingsÚsynced_gpusrÑ   Údropout_probabilityÚskip_the_layerÚlayer_outputss                  r8   r—   zWav2Vec2Encoder.forwardñ  s[  € ñ #7™B¸DÐÙ$5™b¸4ÐàÐ%à$2×$<Ñ$<¸RÓ$@×$GÑ$GÈÈ1Èm×NaÑNaÐbcÑNdÓ$eÐ!Ø45ˆMÐ0Ð0Ñ1Ø×*Ò*à4BÐ4NÐSTÐXfÑSf¡Ðmq‘ð "% ~²a¸¸tÂQÐ6FÑ'G×'JÑ'JÐQ^×QdÑQdÐ'JÓ'eÑ!e�Ø!/´%·+±+¸m×>QÑ>QÓ2R×2VÑ2VÑ!V�Ø!/×!6Ñ!6Ø"×(Ñ(¨Ñ+¨Q°×0DÑ0DÀRÑ0HÈ.×J^ÑJ^Ð_aÑJbó"�ð #×1Ñ1°-Ó@ÐØ%Ð(;Ñ;ˆØŸ™¨Ó6ˆØŸ™ ]Ó3ˆä0Ó2ÒRÔ6LÈTÓ6Rˆà—[‘[ò 	PˆEÙ#Ø$5¸Ð8HÑ$HÐ!ô #(§*¡*¨R£.Ðà%)§]¢]Ð8KÈdÏkÉk×NcÑNcÒ8c™TÐjoˆNÙ!¡[à×.Ò.°4·=²=Ø$(×$EÑ$EØŸ™Ø%Ø&Ø)ó	%‘Mñ %*Ø%°nÐXiô%�Mð !.¨aÑ 0�áØ ,�â Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð7	Pñ:  Ø 1°]Ð4DÑ DÐáÜÑm ]Ð4EÐGZÐ$[ÔmÓmÐmÜØ+Ø+Ø*ô
ð 	
r7   ©NFFT)r/   r0   r1   r„   r3   r  r   r,  rW   r—   rš   r›   s   @r8   ru  ru  æ  s_   ø„ ôYð 26Ø"'Ø%*Ø ñG
à—|‘|ðG
ð ! §¡Ñ.ðG
ð  ð	G
ð
 #ðG
ð ÷G
r7   ru  c                   ó.   ‡ — e Zd Zˆ fd„Z	 	 	 	 dd„Zˆ xZS )ÚWav2Vec2EncoderStableLayerNormc                 óÀ  •— t         ‰| �  «        || _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _	        t        j                  |j                  «      | _        t        j                  t        |j                  «      D �cg c]  }t!        |«      ‘Œ c}«      | _        d| _        |j&                  dk(  | _        y c c}w rw  )rƒ   r„   r‘   r®   rx  r   r    r¸   ró   r¡   rö   rX  rø   rÔ   rU   ry  rm  rz  rÖ   r`  r{  r|  s      €r8   r„   z'Wav2Vec2EncoderStableLayerNorm.__init__<  s©   ø€ Ü‰ÑÔØˆŒÜ=¸fÓEˆÔÜŸ,™, v×'9Ñ'9¸v×?TÑ?TÔUˆŒÜ—z‘z &×"7Ñ"7Ó8ˆŒÜ—m‘mÜBGÈ×H`ÑH`ÓBaÖb¸QÔ0°Õ8Òbó
ˆŒð ',ˆÔ#Ø&,×&AÑ&AÐEXÑ&XˆÕ#ùò cr}  c                 óf  — |rdnd }|rdnd }|�ö|j                  d«      j                  dd|j                  d   «      }||j                  |j                  ¬«      z  }| j
                  r|�d|v r|nd }n‘d|d d …d d d d …f   j                  |j                  ¬«      z
  }|t        j                  |j                  «      j                  z  }|j                  |j                  d   d|j                  d   |j                  d   «      }| j                  |«      }	||	z   }| j                  |«      }t        «       xs t        | «      }
| j                  D ]£  }|r||fz   }t        j                  g «      }| j                   r|| j"                  j$                  k  rdnd	}|r|
rG| j&                  r+| j                   r| j)                  |j*                  |||«      }n ||||¬
«      }|d   }|rd}|sŒ›|d   fz   }Œ¥ | j-                  |«      }|r||fz   }|st/        d„ |||fD «       «      S t1        |||¬«      S )Nr6   rI   r    r$   rJ   r   r�  TFrh  r‚  c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wr•   r6   r„  s     r8   r‡  z9Wav2Vec2EncoderStableLayerNorm.forward.<locals>.<genexpr>Œ  rˆ  r‰  rŠ  )rŒ  r�  r9   r>  rK   r{  r3   rŽ  r�  r�  rx  rø   r   r   rz  rP   rá   r‘   r‘  rÖ   râ   rã   r¡   r’  r   r“  s                  r8   r—   z&Wav2Vec2EncoderStableLayerNorm.forwardH  sn  € ñ #7™B¸DÐÙ$5™b¸4ÐàÐ%à$2×$<Ñ$<¸RÓ$@×$GÑ$GÈÈ1Èm×NaÑNaÐbcÑNdÓ$eÐ!Ø)Ð,A×,DÑ,DÈ=×K^ÑK^Ð,DÓ,_Ñ_ˆMØ×*Ò*à4BÐ4NÐSTÐXfÑSf¡Ðmq‘ð "% ~²a¸¸tÂQÐ6FÑ'G×'JÑ'JÐQ^×QdÑQdÐ'JÓ'eÑ!e�Ø!/´%·+±+¸m×>QÑ>QÓ2R×2VÑ2VÑ!V�Ø!/×!6Ñ!6Ø"×(Ñ(¨Ñ+¨Q°×0DÑ0DÀRÑ0HÈ.×J^ÑJ^Ð_aÑJbó"�ð #×1Ñ1°-Ó@ÐØ%Ð(;Ñ;ˆØŸ™ ]Ó3ˆä0Ó2ÒRÔ6LÈTÓ6Rˆà—[‘[ò 	PˆEÙ#Ø$5¸Ð8HÑ$HÐ!ô #(§*¡*¨R£.Ðà%)§]¢]Ð8KÈdÏkÉk×NcÑNcÒ8c™TÐjoˆNÙ!¡[ð ×.Ò.°4·=²=Ø$(×$EÑ$EØŸ™Ø%Ø&Ø)ó	%‘Mñ %*Ø%°nÐXiô%�Mð !.¨aÑ 0�áØ ,�â Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð9	Pð< Ÿ™¨Ó6ˆáØ 1°]Ð4DÑ DÐáÜÑm ]Ð4EÐGZÐ$[ÔmÓmÐmÜØ+Ø+Ø*ô
ð 	
r7   rœ  r™   r›   s   @r8   rž  rž  ;  s   ø„ ô
Yð ØØ"Ø÷I
r7   rž  c                   ó<   ‡ — e Zd ZdZˆ fd„Zedd„«       Zdd„Zˆ xZS )ÚWav2Vec2GumbelVectorQuantizerz­
    Vector quantization using gumbel softmax. See `[CATEGORICAL REPARAMETERIZATION WITH
    GUMBEL-SOFTMAX](https://arxiv.org/pdf/1611.01144.pdf) for more information.
    c                 ó0  •— t         ‰| �  «        |j                  | _        |j                  | _        |j                  | j                  z  dk7  r&t        d|j                  › d| j                  › d�«      ‚t        j                  t        j                  d| j                  | j
                  z  |j                  | j                  z  «      «      | _        t        j                  |j                  d   | j                  | j
                  z  «      | _        d| _        y )Nr   z`config.codevector_dim z5 must be divisible by `config.num_codevector_groups` z for concatenationr    rI   r$   )rƒ   r„   Únum_codevector_groupsr¨   Únum_codevectors_per_groupÚnum_varsÚcodevector_dimrM   r   Ú	Parameterr3   r4   Úcodevectorsrô   r…   Úweight_projÚtemperaturerí   s     €r8   r„   z&Wav2Vec2GumbelVectorQuantizer.__init__š  sì   ø€ Ü‰ÑÔØ ×6Ñ6ˆŒØ×8Ñ8ˆŒà× Ñ  4§?¡?Ñ2°aÒ7ÜØ)¨&×*?Ñ*?Ð)@ð A5Ø59·_±_Ð4EÐEWðYóð ô Ÿ<™<Ü×Ñ˜a §¡°4·=±=Ñ!@À&×BWÑBWÐ[_×[jÑ[jÑBjÓkó
ˆÔô Ÿ9™9 V§_¡_°RÑ%8¸$¿/¹/ÈDÏMÉMÑ:YÓZˆÔð ˆÕr7   c           	      óÐ  — |�|j                  «       d d …d d f   j                  | j                  «      }t        j                  || t        j
                  | «      «      } | j                  d¬«      |j                  «       z  }n| j                  d¬«      }t        j                  t        j                  |t        j                  |dz   «      z  d¬«       «      j                  «       }|S )Nr   r  gH¯¼šò×z>rI   )
Úflattenr�  r9   r3   ÚwhereÚ
zeros_likerS   ÚmeanÚexpÚlog)ÚprobsÚmaskÚmask_extendedÚmarginal_probsÚ
perplexitys        r8   Ú_compute_perplexityz1Wav2Vec2GumbelVectorQuantizer._compute_perplexity®  sµ   € àÐØ ŸL™L›Nª1¨d°D¨=Ñ9×@Ñ@ÀÇÁÓMˆMÜ—K‘K ¨u´e×6FÑ6FÀuÓ6MÓNˆEØ"ŸY™Y¨1˜YÓ-°·±³
Ñ:‰Nà"ŸZ™Z¨A˜ZÓ.ˆNä—Y‘Y¤§	¡	¨.¼5¿9¹9À^ÐVZÑEZÓ;[Ñ*[ÐacÔ dÐdÓe×iÑiÓkˆ
ØÐr7   c                 óæ  — |j                   \  }}}| j                  |«      }|j                  ||z  | j                  z  d«      }| j                  rŸt
        j                  j                  |j                  «       | j                  d¬«      j                  |«      }t        j                  |j                  ||z  | j                  d«      j                  «       d¬«      }| j                  ||«      }n€|j                  d¬«      }	|j                  |j                   «      j!                  d|	j                  dd«      d«      }|j                  ||z  | j                  d«      }| j                  ||«      }|j                  ||z  d«      }|j#                  d«      | j$                  z  }
|
j                  ||z  | j                  | j&                  d«      }|j)                  d«      j                  ||d«      }||fS )NrI   T)ÚtauÚhardr  r    r�  r£   )r9   r«  r  r¨   rá   r   r  Úgumbel_softmaxr+  r¬  Útype_asr3   r  r¹  ÚargmaxÚ	new_zerosÚscatter_rŒ  rª  r§  rS   )r�   r+   ro   rc   rG   r¸   Úcodevector_probsÚcodevector_soft_distr¸  Úcodevector_idxÚcodevectors_per_grouprª  s               r8   r—   z%Wav2Vec2GumbelVectorQuantizer.forwardº  sß  € Ø3@×3FÑ3FÑ0ˆ
�O [ð ×(Ñ(¨Ó7ˆØ%×*Ñ*¨:¸Ñ+GÈ$Ï/É/Ñ+YÐ[]Ó^ˆà�=Š=ä!Ÿ}™}×;Ñ;Ø×#Ñ#Ó%¨4×+;Ñ+;À$ð  <ó  ç‰g�mÓ$ð ô
 $)§=¡=Ø×"Ñ" :°Ñ#?ÀÇÁÐRTÓU×[Ñ[Ó]Ðceô$Ð ð ×1Ñ1Ð2FÐHYÓZ‰Jð +×1Ñ1°bÐ1Ó9ˆNØ,×6Ñ6°}×7JÑ7JÓK×TÑTØ�N×'Ñ'¨¨AÓ.°ó Ðð  0×4Ñ4°ZÀ/Ñ5QÐSW×SbÑSbÐdfÓgÐà×1Ñ1Ð2BÐDUÓVˆJà+×0Ñ0°¸oÑ1MÈrÓRÐà 0× :Ñ :¸2Ó >À×AQÑAQÑ QÐØ+×0Ñ0°¸oÑ1MÈtÏÉÐ`d×`mÑ`mÐoqÓrˆØ!—o‘o bÓ)×.Ñ.¨z¸?ÈBÓOˆà˜JÐ&Ð&r7   r•   )	r/   r0   r1   r2   r„   Ústaticmethodr¹  r—   rš   r›   s   @r8   r£  r£  ”  s&   ø„ ñô
ð( ò	ó ð	÷#'r7   r£  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚWav2Vec2Adapterc                 ó¨  •‡— t         ‰| �  «        ‰j                  ‰j                  k7  rTt	        j
                  ‰j                  ‰j                  «      | _        t	        j                  ‰j                  «      | _        nd x| _        | _        t	        j                  ˆfd„t        ‰j                  «      D «       «      | _        ‰j                  | _        y )Nc              3   ó4   •K  — | ]  }t        ‰«      –— Œ y ­wr•   )ÚWav2Vec2AdapterLayer)r…  rd   r‘   s     €r8   r‡  z+Wav2Vec2Adapter.__init__.<locals>.<genexpr>ë  s   øè ø€ Ò#kÀQÔ$8¸×$@Ñ#kùs   ƒ)rƒ   r„   Úoutput_hidden_sizer¸   r   rô   Úprojr    Úproj_layer_normrÔ   rU   Únum_adapter_layersrz  r‘  rí   s    `€r8   r„   zWav2Vec2Adapter.__init__á  s—   ù€ Ü‰ÑÔð ×$Ñ$¨×(:Ñ(:Ò:ÜŸ	™	 &×"4Ñ"4°f×6OÑ6OÓPˆDŒIÜ#%§<¡<°×0IÑ0IÓ#JˆDÕ à/3Ð3ˆDŒI˜Ô,ä—m‘mÓ#kÌ%ÐPV×PiÑPiÓJjÔ#kÓkˆŒØ×)Ñ)ˆ�r7   c                 óh  — | j                   �.| j                  �"| j                  |«      }| j                  |«      }|j                  dd«      }| j                  D ]D  }t        j
                  j                  «       }| j                  r|| j                  kD  sŒ= ||«      }ŒF |j                  dd«      }|S rÇ   )rÍ  rÎ  r¤   rz  rN   rO   rá   r‘  )r�   r+   rÑ   Úlayerdrop_probs       r8   r—   zWav2Vec2Adapter.forwardî  s¢   € à�9‰9Ð  T×%9Ñ%9Ð%EØ ŸI™I mÓ4ˆMØ ×0Ñ0°Ó?ˆMà%×/Ñ/°°1Ó5ˆà—[‘[ò 	5ˆEÜŸY™Y×-Ñ-Ó/ˆNØ—=’= ^°d·n±nÓ%DÙ % mÓ 4‘ð	5ð
 &×/Ñ/°°1Ó5ˆØÐr7   r™   r›   s   @r8   rÈ  rÈ  à  s   ø„ ô*ör7   rÈ  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )rË  c                 ó¶   •— t         ‰| �  «        t        j                  |j                  d|j                  z  |j
                  |j                  d¬«      | _        y )Nr$   r    )r�   r°   )rƒ   r„   r   rˆ   rÌ  Úadapter_kernel_sizeÚadapter_striderŒ   rí   s     €r8   r„   zWav2Vec2AdapterLayer.__init__   sJ   ø€ Ü‰ÑÔÜ—I‘IØ×%Ñ%Ø�×)Ñ)Ñ)Ø×&Ñ&Ø×(Ñ(Øô
ˆ�	r7   c                 ój   — | j                  |«      }t        j                  j                  |d¬«      }|S )Nr    r  )rŒ   r   r  Úglur–   s     r8   r—   zWav2Vec2AdapterLayer.forward
  s/   € ØŸ	™	 -Ó0ˆÜŸ™×)Ñ)¨-¸QÐ)Ó?ˆàÐr7   r™   r›   s   @r8   rË  rË  ÿ  s   ø„ ô
ör7   rË  c                   ó>   ‡ — e Zd Zˆ fd„Zdej
                  fd„Zˆ xZS )rq  c                 óœ  •— t         ‰| �  «        |j                  | _        |j                  | _        t        j                  | j
                  «      | _        t        j                  | j
                  | j                  «      | _
        t        j                  «       | _        t        j                  | j                  | j
                  «      | _        y)zŸ
        Implements adapter modules directly with 3D tensor weight as parameters and without using ModuleList to speed
        up training throughput.
        N)rƒ   r„   ro  Ú	input_dimr¸   Ú
hidden_dimr   r    Únormrô   Úlinear_1ÚReLUÚact_fnÚlinear_2rí   s     €r8   r„   z!Wav2Vec2AttnAdapterLayer.__init__  s   ø€ ô
 	‰ÑÔØ×0Ñ0ˆŒØ ×,Ñ,ˆŒä—L‘L §¡Ó1ˆŒ	ÜŸ	™	 $§/¡/°4·>±>ÓBˆŒÜ—g‘g“iˆŒÜŸ	™	 $§.¡.°$·/±/ÓBˆ�r7   r+   c                 óŽ   — | j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }|S r•   )rÜ  rÝ  rß  rà  r–   s     r8   r—   z Wav2Vec2AttnAdapterLayer.forward   s@   € ØŸ	™	 -Ó0ˆàŸ™ mÓ4ˆØŸ™ MÓ2ˆØŸ™ mÓ4ˆàÐr7   )r/   r0   r1   r„   r3   r4   r—   rš   r›   s   @r8   rq  rq    s   ø„ ôCð U×%6Ñ%6÷ r7   rq  c                   ó¨   — e Zd ZdZeZdZdZdZdZ	dZ
d„ Z	 ddeej                  ef   dee   fd	„Z	 dd
edej                  fd„Zd„ Zd„ Zddefd„Zy)ÚWav2Vec2PreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úwav2vec2rä   Tc           
      óH  — t        |t        «      rW|j                  j                  «        |j                  j                  «        d|j                  _        d|j                  _        yt        |t        «      r‰|j                  j                  j                  j                  dd¬«       |j                  j                  j                  j                  «        t        j                  j                  |j                   «       yt        |t"        «      r²t        j                  j                  |j$                  j                  ddt'        j(                  d|j$                  j*                  d   |j$                  j,                  z  z  «      z  ¬«       t        j                  j/                  |j$                  j                  d«       yt        |t0        «      r›t'        j(                  d|j2                  j4                  z  «      }t        j                  j                  |j2                  j                  | |¬«       t        j                  j                  |j2                  j                  | |¬«       yt        |t        j6                  «      rm|j                  j                  j                  d| j8                  j:                  ¬«       |j                  �%|j                  j                  j                  «        yyt        |t        j<                  t        j>                  f«      rJ|j                  j                  j                  «        |j                  j                  jA                  d	«       yt        |t        jB                  «      r t        j                  jE                  |j                  «       |j                  �jt'        j(                  |jF                  |j,                  |j*                  d   z  z  «      }t        j                  j                  |j                  | |¬«       yyy)
zInitialize the weightsTr)  r    )r±  Ústdr   r$   )ÚaÚbNr�  )$rS  ÚWav2Vec2ForPreTrainingÚproject_hidÚreset_parametersÚ	project_qÚ_is_hf_initializedr£  r«  r´   ÚdataÚnormal_r‚   Úzero_r   ÚinitÚuniform_rª  r®   rŒ   ÚmathÚsqrtr€   Úin_channelsÚ	constant_rï   rõ   Úin_featuresrô   r‘   Úinitializer_ranger    r«   Úfill_rˆ   Úkaiming_normal_r±   )r�   ÚmoduleÚks      r8   Ú_init_weightsz%Wav2Vec2PreTrainedModel._init_weights7  sÕ  € ô �fÔ4Ô5Ø×Ñ×/Ñ/Ô1Ø×Ñ×-Ñ-Ô/Ø48ˆF×ÑÔ1Ø26ˆF×ÑÕ/ä˜Ô =Ô>Ø×Ñ×%Ñ%×*Ñ*×2Ñ2¸ÀÐ2ÔCØ×Ñ×#Ñ#×(Ñ(×.Ñ.Ô0Ü�G‰G×Ñ˜V×/Ñ/Õ0Ü˜Ô ?Ô@Ü�G‰G�O‰OØ—‘×"Ñ"ØØœŸ	™	 ! v§{¡{×'>Ñ'>¸qÑ'AÀFÇKÁK×D[ÑD[Ñ'[Ñ"\Ó]Ñ]ð ô ô
 �G‰G×Ñ˜fŸk™k×.Ñ.°Õ2Ü˜Ô 9Ô:Ü—	‘	˜!˜f×/Ñ/×;Ñ;Ñ;Ó<ˆAÜ�G‰G×Ñ˜V×.Ñ.×5Ñ5¸!¸¸qÐÔAÜ�G‰G×Ñ˜V×.Ñ.×3Ñ3¸°r¸QÐÕ?Ü˜¤§	¡	Ô*Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSà�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡¬r¯|©|Ð <Ô=Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)Ü˜¤§	¡	Ô*Ü�G‰G×#Ñ# F§M¡MÔ2à�{‰{Ð&Ü—I‘I˜fŸm™m¨v×/AÑ/AÀF×DVÑDVÐWXÑDYÑ/YÑZÓ[�Ü—‘× Ñ  §¡°°°aÐ Õ8ð 'ð +r7   Nre   Úadd_adapterc                 óT  — |€| j                   j                  n|}d„ }t        | j                   j                  | j                   j                  «      D ]  \  }} ||||«      }Œ |rBt        | j                   j                  «      D ]   } ||d| j                   j                  «      }Œ" |S )zH
        Computes the output length of the convolutional layers
        c                 ó>   — t        j                  | |z
  |d¬«      dz   S )NÚfloor)Úrounding_moder    )r3   Údiv©rD   r€   r�   s      r8   Ú_conv_out_lengthzRWav2Vec2PreTrainedModel._get_feat_extract_output_lengths.<locals>._conv_out_lengthg  s"   € ô —9‘9˜\¨KÑ7¸ÈwÔWÐZ[Ñ[Ð[r7   r    )r‘   rþ  Úzipr‰   rŠ   rU   rÏ  rÕ  )r�   re   rþ  r  r€   r�   rd   s          r8   Ú _get_feat_extract_output_lengthsz8Wav2Vec2PreTrainedModel._get_feat_extract_output_lengths^  s¦   € ð 2=Ð1D�d—k‘k×-Ò-È+ˆò	\ô
 $' t§{¡{×'>Ñ'>ÀÇÁ×@WÑ@WÓ#Xò 	QÑˆK˜Ù,¨]¸KÈÓP‰Mð	Qñ Ü˜4Ÿ;™;×9Ñ9Ó:ò _�Ù 0°ÀÀ4Ç;Á;×C]ÑC]Ó ^‘ð_ð Ðr7   Úfeature_vector_lengthr<   c                 ó   — |j                  d¬«      d d …df   }| j                  ||¬«      }|j                  t        j                  «      }|j
                  d   }t        j                  ||f|j                  |j                  ¬«      }d|t        j                  |j
                  d   |j                  ¬«      |dz
  f<   |j                  dg«      j                  d«      j                  dg«      j                  «       }|S )NrI   r  ©rþ  r   )rK   Údevicer    )r  )Úcumsumr  r>  r3   Úlongr9   rV   rK   r  rY   ÚfliprW   )r�   r  r<   rþ  Únon_padded_lengthsÚoutput_lengthsrc   s          r8   Ú"_get_feature_vector_attention_maskz:Wav2Vec2PreTrainedModel._get_feature_vector_attention_masku  só   € ð
 ,×2Ñ2°rÐ2Ó:º1¸b¸5ÑAÐà×>Ñ>Ð?QÐ_jÐ>ÓkˆØ'×*Ñ*¬5¯:©:Ó6ˆà#×)Ñ)¨!Ñ,ˆ
äŸ™ØÐ.Ð/°~×7KÑ7KÐTb×TiÑTiô
ˆð uvˆœŸ™ ^×%9Ñ%9¸!Ñ%<À^×EZÑEZÔ[Ð]kÐnoÑ]oÐpÑqØ'×,Ñ,¨b¨TÓ2×9Ñ9¸"Ó=×BÑBÀBÀ4ÓH×MÑMÓOˆØÐr7   c                 ó¤  — | j                   j                  €t        | j                  › d�«      ‚i }| j	                  «       D ]D  \  }}t        |t        «      sŒ|j                  «       D ]  \  }}||dj                  ||g«      <   Œ ŒF t        | t        «      r8| j                  j                  «       D ]  \  }}||dj                  d|g«      <   Œ |S )NzF has no adapter layers. Make sure to define `config.adapter_attn_dim`.r7  Úlm_head)r‘   ro  rM   r“   Únamed_modulesrS  rq  Únamed_parametersÚjoinÚWav2Vec2ForCTCr  )r�   Úadapter_weightsrµ   rû  Ú
param_namerÞ   s         r8   Ú_get_adaptersz%Wav2Vec2PreTrainedModel._get_adapters‰  sÝ   € Ø�;‰;×'Ñ'Ð/Ü §¡Ð/Ð/uÐvÓwÐwàˆØ ×.Ñ.Ó0ò 	J‰LˆD�&Ü˜&Ô":Õ;Ø)/×)@Ñ)@Ó)Bò JÑ%�J ØDI�O C§H¡H¨d°JÐ-?Ó$@ÒAñJð	Jô
 �dœNÔ+Ø#Ÿ|™|×<Ñ<Ó>ò E‘��eØ?D� §¡¨)°TÐ):Ó ;Ò<ðEð Ðr7   c                 óÊ   — | j                  «       D ]$  }t        |t        «      sŒ| j                  |«       Œ& t        | t        «      r| j                  | j
                  «       yy)zc
        (Re-)initialize attention adapter layers and lm head for adapter-only fine-tuning
        N)ÚmodulesrS  rq  rý  r  r  )r�   rû  s     r8   Úinit_adapter_layersz+Wav2Vec2PreTrainedModel.init_adapter_layers™  sU   € ð
 —l‘l“nò 	+ˆFÜ˜&Ô":Õ;Ø×"Ñ" 6Õ*ð	+ô
 �dœNÔ+Ø×Ñ˜tŸ|™|Õ,ð ,r7   Útarget_langc                 ó–  — | j                   j                  €t        d|› d�«      ‚|| j                  k(  r|st        j                  d|› d�«       y|j                  dd«      }|j                  dd«      }|j                  d	d«      }|j                  d
d«      }|j                  dd«      }|j                  dd«      }	|j                  dd«      }
|j                  dd«      }|j                  dt        «       rdnd«      }|
�)t        j                  dt        «       |	�t        d«      ‚|
}	| j                   j                  }d}|dur5t        j                  |«      }	 t        |||||||	||¬«	      }t        |«      }|€Bt$        j                  |«      }	 t        |||||||	||¬«	      }t'        j(                  |dd¬«      }| j+                  «       }t-        |j/                  «       «      t-        |j/                  «       «      z
  }t-        |j/                  «       «      t-        |j/                  «       «      z
  }t1        |«      dkD  r!t        d› ddj3                  |«      › d�«      ‚t1        |«      dkD  r!t        d› ddj3                  |«      › d�«      ‚|d   j4                  d   }|| j                   j6                  k7  rWt9        j:                  | j                   j<                  || j>                  | j@                  ¬«      | _!        || j                   _        |jE                  «       D ��ci c]  \  }}||jG                  ||   «      “Œ }}}| jI                  |d¬ «       || _        y# t         $ r |r‚ Y �Œùt"        $ r |rt!        d|› d|› d|› d�«      ‚Y �Œw xY w# t         $ r ‚ t"        $ r t!        d|› d|› d|› d�«      ‚w xY wc c}}w )!aÇ  
        Load a language adapter model from a pre-trained adapter model.

        Parameters:
            target_lang (`str`):
                Has to be a language id of an existing adapter weight. Adapter weights are stored in the format
                adapter.<lang>.safetensors or adapter.<lang>.bin
            force_load (`bool`, defaults to `True`):
                Whether the weights shall be loaded even if `target_lang` matches `self.target_lang`.
            cache_dir (`Union[str, os.PathLike]`, *optional*):
                Path to a directory in which a downloaded pretrained model configuration should be cached if the
                standard cache should not be used.
            force_download (`bool`, *optional*, defaults to `False`):
                Whether or not to force the (re-)download of the model weights and configuration files, overriding the
                cached versions if they exist.
            resume_download:
                Deprecated and ignored. All downloads are now resumed by default when possible.
                Will be removed in v5 of Transformers.
            proxies (`Dict[str, str]`, *optional*):
                A dictionary of proxy servers to use by protocol or endpoint, e.g., `{'http': 'foo.bar:3128',
                'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request.
            local_files_only(`bool`, *optional*, defaults to `False`):
                Whether or not to only look at local files (i.e., do not try to download the model).
            token (`str` or `bool`, *optional*):
                The token to use as HTTP bearer authorization for remote files. If `True`, or not specified, will use
                the token generated when running `huggingface-cli login` (stored in `~/.huggingface`).
            revision (`str`, *optional*, defaults to `"main"`):
                The specific model version to use. It can be a branch name, a tag name, or a commit id, since we use a
                git-based system for storing models and other artifacts on huggingface.co, so `revision` can be any
                identifier allowed by git.

                <Tip>

                To test a pull request you made on the Hub, you can pass `revision="refs/pr/<pr_number>"`.

                </Tip>

            mirror (`str`, *optional*):
                Mirror source to accelerate downloads in China. If you are from China and have an accessibility
                problem, you can set this option to resolve it. Note that we do not guarantee the timeliness or safety.
                Please refer to the mirror site for more information.

        <Tip>

        Activate the special ["offline-mode"](https://huggingface.co/transformers/installation.html#offline-mode) to
        use this method in a firewalled environment.

        </Tip>

        Examples:

        ```python
        >>> from transformers import Wav2Vec2ForCTC, AutoProcessor

        >>> ckpt = "facebook/mms-1b-all"
        >>> processor = AutoProcessor.from_pretrained(ckpt)
        >>> model = Wav2Vec2ForCTC.from_pretrained(ckpt, target_lang="eng")
        >>> # set specific language
        >>> processor.tokenizer.set_target_lang("spa")
        >>> model.load_adapter("spa")
        ```
        NzCannot load_adapter for ú- if `config.adapter_attn_dim` is not defined.z#Adapter weights are already set to r7  Ú	cache_dirÚforce_downloadFÚresume_downloadÚproxiesÚlocal_files_onlyÚtokenÚuse_auth_tokenÚrevisionÚuse_safetensorszrThe `use_auth_token` argument is deprecated and will be removed in v5 of Transformers. Please use `token` instead.zV`token` and `use_auth_token` are both specified. Please set only the argument `token`.)Úfilenamer"  r#  r$  r%  r&  r(  r!  zCan't load the model for 'zœ'. If you were trying to load it from 'https://huggingface.co/models', make sure you don't have a local directory with the same name. Otherwise, make sure 'z=' is the correct path to a directory containing a file named ÚcpuT)Úmap_locationÚweights_onlyr   zThe adapter weights z has unexpected keys: z, z has missing keys: zlm_head.weight©r  rK   )Ústrict)%r‘   ro  rM   r  r<  ÚwarningÚpopr   ré   rê   rì   Ú_name_or_pathÚWAV2VEC2_ADAPTER_SAFE_FILEÚformatr   Úsafe_load_fileÚEnvironmentErrorÚ	ExceptionÚWAV2VEC2_ADAPTER_PT_FILEr3   Úloadr  ÚsetÚkeysrZ   r  r9   Ú
vocab_sizer   rô   rÌ  r  rK   r  Úitemsr>  Úload_state_dict)r�   r  Ú
force_loadr2  r!  r"  r#  r$  r%  r&  r'  r(  r)  Úmodel_path_or_idÚ
state_dictÚfilepathÚweight_pathr  Úunexpected_keysÚmissing_keysÚtarget_vocab_sizerü  r†  s                          r8   Úload_adapterz$Wav2Vec2PreTrainedModel.load_adapter¦  s.  € ð~ �;‰;×'Ñ'Ð/ÜÐ7¸°}ÐDqÐrÓsÐsà˜$×*Ñ*Ò*±:Ü�N‰NÐ@ÀÀÈQÐOÔPØà—J‘J˜{¨DÓ1ˆ	ØŸ™Ð$4°eÓ<ˆØ Ÿ*™*Ð%6¸Ó=ˆØ—*‘*˜Y¨Ó-ˆØ!Ÿ:™:Ð&8¸%Ó@ÐØ—
‘
˜7 DÓ)ˆØŸ™Ð$4°dÓ;ˆØ—:‘:˜j¨$Ó/ˆØ Ÿ*™*Ð%6Ô@XÔ@Z¹Ð`eÓfˆàÐ%Ü�M‰Mð EÜôð Ð Ü Ølóð ð #ˆEàŸ;™;×4Ñ4ÐØˆ
ð  %Ñ'Ü1×8Ñ8¸ÓEˆHðÜ)Ø$Ø%Ø#1Ø$3Ø#Ø%5ØØ%Ø'ô
�ô ,¨KÓ8�
ð& ÐÜ/×6Ñ6°{ÓCˆHðÜ)Ø$Ø%Ø#1Ø$3Ø#Ø%5ØØ%Ø'ô
�ô #ŸZ™ZØØ!&Ø!%ô�
ð( ×,Ñ,Ó.ˆÜ˜jŸo™oÓ/Ó0´3°×7KÑ7KÓ7MÓ3NÑNˆÜ˜?×/Ñ/Ó1Ó2´S¸¿¹Ó9JÓ5KÑKˆäˆÓ !Ò#ÜÐ3°K°=Ð@VÐW[×W`ÑW`ÐapÓWqÐVrÐrsÐtÓuÐuÜ�Ó Ò"ÜÐ3°K°=Ð@SÐTX×T]ÑT]Ð^jÓTkÐSlÐlmÐnÓoÐoð 'Ð'7Ñ8×>Ñ>¸qÑAÐØ §¡× 6Ñ 6Ò6ÜŸ9™9Ø—‘×.Ñ.Ð0AÈ$Ï+É+Ð]a×]gÑ]gôˆDŒLð &7ˆD�K‰KÔ"ð ?I×>NÑ>NÓ>P×Q±d°a¸�a˜Ÿ™˜o¨aÑ0Ó1Ñ1ÐQˆ
ÑQØ×Ñ˜Z°ÐÔ6ð 'ˆÕøôW $ò Ù"ð ò #ô
 ò á"Ü*Ø4Ð5EÐ4Fð G=à=MÐ<Nð O>Ø>F¸ZÀqðJóð ò #ðûôB $ò ð äò ä&Ø0Ð1AÐ0Bð C9à9IÐ8Jð K:Ø:B¸À1ðFóð ðüó6 Rs*   ÅM% Æ,N Ì(OÍ%NÍ4NÎNÎ(Or•   )T)r/   r0   r1   r2   r!   Úconfig_classÚbase_model_prefixÚmain_input_nameÚsupports_gradient_checkpointingÚ_supports_flash_attn_2Ú_supports_sdparý  r   r3   Ú
LongTensorrB   r   rW   r  r  r  r  rU  rG  r6   r7   r8   rã  rã  *  s›   „ ñð
 "€LØ"ÐØ$€OØ&*Ð#Ø!ÐØ€Nò%9ðP Z^ñØ" 5×#3Ñ#3°SÐ#8Ñ9ðØHPÐQUÉóð0 Y]ñØ%(ðØ:?×:JÑ:Jóò(ò -ñ|'¨ô |'r7   rã  aè  
    Wav2Vec2 was proposed in [wav2vec 2.0: A Framework for Self-Supervised Learning of Speech
    Representations](https://arxiv.org/abs/2006.11477) by Alexei Baevski, Henry Zhou, Abdelrahman Mohamed, Michael
    Auli.

    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
    library implements for all its model (such as downloading or saving etc.).

    This model is a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) sub-class. Use
    it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage and
    behavior.

    Parameters:
        config ([`Wav2Vec2Config`]): Model configuration class with all the parameters of the model.
            Initializing with a config file does not load the weights associated with the model, only the
            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
a¢  
    Args:
        input_values (`torch.FloatTensor` of shape `(batch_size, sequence_length)`):
            Float values of input raw speech waveform. Values can be obtained by loading a `.flac` or `.wav` audio file
            into an array of type `List[float]` or a `numpy.ndarray`, *e.g.* via the soundfile library (`pip install
            soundfile`). To prepare the array into `input_values`, the [`AutoProcessor`] should be used for padding and
            conversion into a tensor of type `torch.FloatTensor`. See [`Wav2Vec2Processor.__call__`] for details.
        attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to avoid performing convolution and attention on padding token indices. Mask values selected in `[0,
            1]`:

            - 1 for tokens that are **not masked**,
            - 0 for tokens that are **masked**.

            [What are attention masks?](../glossary#attention-mask)

            <Tip warning={true}>

            `attention_mask` should only be passed if the corresponding processor has `config.return_attention_mask ==
            True`. For all models whose processor has `config.return_attention_mask == False`, such as
            [wav2vec2-base](https://huggingface.co/facebook/wav2vec2-base-960h), `attention_mask` should **not** be
            passed to avoid degraded performance when doing batched inference. For such models `input_values` should
            simply be padded with 0 and passed without `attention_mask`. Be aware that these models also yield slightly
            different results depending on whether `input_values` is padded or not.

            </Tip>

        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
zbThe bare Wav2Vec2 Model transformer outputting raw hidden-states without any specific head on top.c                   ób  ‡ — e Zd Zdefˆ fd„Zd„ Zd„ Z	 	 ddej                  de	ej                     de	ej                     fd„Z ee«       eeeed	e¬
«      	 	 	 	 	 dde	ej&                     de	ej&                     de	ej                     de	e   de	e   de	e   deeef   fd„«       «       Zˆ xZS )ÚWav2Vec2Modelr‘   c                 óî  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        |j                  dkD  s|j                  dkD  rEt        j                  t        j                  |j                  «      j                  «       «      | _        |j                   rt#        |«      | _        nt'        |«      | _        |j(                  rt+        |«      nd | _        | j/                  «        y )Nr)  )rƒ   r„   r‘   rÎ   Úfeature_extractorrï   Úfeature_projectionÚmask_time_probÚmask_feature_probr   r©  r3   r,  r¸   rò  Úmasked_spec_embedÚdo_stable_layer_normrž  Úencoderru  rþ  rÈ  ÚadapterÚ	post_initrí   s     €r8   r„   zWav2Vec2Model.__init__£  sº   ø€ Ü‰Ñ˜Ô ØˆŒÜ!7¸Ó!?ˆÔÜ";¸FÓ"CˆÔð × Ñ  3Ò&¨&×*BÑ*BÀSÒ*HÜ%'§\¡\´%·,±,¸v×?QÑ?QÓ2R×2[Ñ2[Ó2]Ó%^ˆDÔ"à×&Ò&Ü9¸&ÓAˆD�Lä*¨6Ó2ˆDŒLà28×2DÒ2D” vÔ.È$ˆŒð 	�‰Õr7   c                 óX   — t        j                  dt        «       | j                  «        y©z©
        Calling this function will disable the gradient computation for the feature encoder so that its parameters will
        not be updated during training.
        úžThe method `freeze_feature_extractor` is deprecated and will be removed in Transformers v5. Please use the equivalent `freeze_feature_encoder` method instead.N©ré   rê   rì   Úfreeze_feature_encoder©r�   s    r8   Úfreeze_feature_extractorz&Wav2Vec2Model.freeze_feature_extractor·  ó'   € ô
 	�‰ðQäô	
ð
 	×#Ñ#Õ%r7   c                 ó8   — | j                   j                  «        y©ú¨
        Calling this function will disable the gradient computation for the feature encoder so that its parameter will
        not be updated during training.
        N)rR  rß   r`  s    r8   r_  z$Wav2Vec2Model.freeze_feature_encoderÃ  s   € ð
 	×Ñ×1Ñ1Õ3r7   r+   ro   r<   c                 óÎ  — t        | j                  dd«      s|S |j                  «       \  }}}|�)| j                  j	                  |j
                  «      ||<   nË| j                  j                  dkD  r²| j                  r¦t        ||f| j                  j                  | j                  j                  || j                  j                  ¬«      }t        j                  ||j                  t        j                  ¬«      }| j                  j	                  |j
                  «      ||<   | j                  j                  dkD  r¨| j                  rœt        ||f| j                  j                  | j                  j                   | j                  j"                  ¬«      }t        j                  ||j                  t        j                  ¬«      }|dd…df   j%                  d|d«      }d||<   |S )	zš
        Masks extracted features along time axis and/or along feature axis according to
        [SpecAugment](https://arxiv.org/abs/1904.08779).
        Úapply_spec_augmentTNr   )r:   r;   r<   r=   r.  )r:   r;   r=   rI   )rp  r‘   rq   rV  r>  rK   rT  rá   rl   Úmask_time_lengthÚmask_time_min_masksr3   r  r  rW   rU  Úmask_feature_lengthÚmask_feature_min_masksr�  )r�   r+   ro   r<   rc   rG   r¸   Úmask_feature_indicess           r8   Ú_mask_hidden_statesz!Wav2Vec2Model._mask_hidden_statesÊ  sš  € ô �t—{‘{Ð$8¸$Ô?Ø Ð ð 4A×3EÑ3EÓ3GÑ0ˆ
�O [àÐ(à/3×/EÑ/E×/HÑ/HÈ×I\ÑI\Ó/]ˆMÐ+Ò,Ø�[‰[×'Ñ'¨!Ò+°·²Ü 5Ø˜_Ð-ØŸ+™+×4Ñ4Ø ŸK™K×8Ñ8Ø-ØŸ+™+×9Ñ9ô!Ðô !&§¡Ð->À}×G[ÑG[Ôch×cmÑcmÔ nÐØ/3×/EÑ/E×/HÑ/HÈ×I\ÑI\Ó/]ˆMÐ+Ñ,à�;‰;×(Ñ(¨1Ò,°·²ä#8Ø˜[Ð)ØŸ+™+×7Ñ7Ø ŸK™K×;Ñ;ØŸ+™+×<Ñ<ô	$Ð ô $)§<¡<Ð0DÈ]×MaÑMaÔin×isÑisÔ#tÐ Ø#7º¸4¸Ñ#@×#GÑ#GÈÈOÐ]_Ó#`Ð Ø23ˆMÐ.Ñ/àÐr7   Úaudio©Ú
checkpointÚoutput_typerH  ÚmodalityÚexpected_outputrä   r  r~  r  r>   c                 óH  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }| j	                  |«      }|j                  dd«      }|�!| j                  |j                  d   |d¬«      }| j                  |«      \  }}| j                  |||¬«      }| j                  |||||¬«      }	|	d   }| j                  �| j                  |«      }|s
||f|	dd  z   S t        |||	j                  |	j                  ¬«      S )	Nr    r$   Fr
  )ro   r<   ©r<   r  r~  r  r   )r‹  Úextract_featuresr+   r,   )r‘   r  r~  Úuse_return_dictrR  r¤   r  r9   rS  rm  rX  rY  r   r+   r,   )
r�   rä   r<   ro   r  r~  r  rv  r+   Úencoder_outputss
             r8   r—   zWav2Vec2Model.forwardø  sb  € ð" 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà×1Ñ1°,Ó?ÐØ+×5Ñ5°a¸Ó;ÐàÐ%à!×DÑDØ ×&Ñ& qÑ)¨>Àuð Eó ˆNð +/×*AÑ*AÐBRÓ*SÑ'ˆÐ'Ø×0Ñ0ØÐ->È~ð 1ó 
ˆð Ÿ,™,ØØ)Ø/Ø!5Ø#ð 'ó 
ˆð (¨Ñ*ˆà�<‰<Ð#Ø ŸL™L¨Ó7ˆMáØ!Ð#3Ð4°ÀqÀrÐ7JÑJÐJä&Ø+Ø-Ø)×7Ñ7Ø&×1Ñ1ô	
ð 	
r7   r‚  ©NNNNN)r/   r0   r1   r!   r„   ra  r_  r3   r4   r   rN  rm  r   ÚWAV2VEC2_INPUTS_DOCSTRINGr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCÚ_EXPECTED_OUTPUT_SHAPEr,  rW   r   r   r—   rš   r›   s   @r8   rP  rP  ž  s  ø„ ð
˜~õ ò(
&ò4ð :>Ø59ñ	,à×(Ñ(ð,ð $ E×$5Ñ$5Ñ6ð,ð ! ×!1Ñ!1Ñ2ó	,ñ\ +Ð+DÓEÙØ&Ø+Ø$ØØ.ôð 26Ø9=Ø,0Ø/3Ø&*ñ2
à˜uŸ|™|Ñ,ð2
ð ! §¡Ñ.ð2
ð $ E×$5Ñ$5Ñ6ð	2
ð
 $ D™>ð2
ð ' t™nð2
ð ˜d‘^ð2
ð 
ˆuÐ-Ð-Ñ	.ò2
óó Fô2
r7   rP  z5Wav2Vec2 Model with a quantizer and `VQ` head on top.c                   óˆ  ‡ — e Zd Zdefˆ fd„Zdefd„Zd„ Zd„ Ze		 dde
j                  de
j                  d	e
j                  defd
„«       Z ee«       eee¬«      	 	 	 	 	 	 ddee
j&                     dee
j&                     dee
j(                     dee
j(                     dee   dee   dee   deeef   fd„«       «       Zˆ xZS )ré  r‘   c                 óˆ  •— t         ‰| �  |«       t        |«      | _        t	        j
                  |j                  «      | _        t        |«      | _	        t	        j                  |j                  |j                  «      | _        t	        j                  |j                  |j                  «      | _        | j!                  «        y r•   )rƒ   r„   rP  rä  r   rö   Úfeat_quantizer_dropoutÚdropout_featuresr£  Ú	quantizerrô   r¸   Úproj_codevector_dimrê  r¨  rì  rZ  rí   s     €r8   r„   zWav2Vec2ForPreTraining.__init__7  sˆ   ø€ Ü‰Ñ˜Ô Ü% fÓ-ˆŒÜ "§
¡
¨6×+HÑ+HÓ IˆÔä6°vÓ>ˆŒäŸ9™9 V×%7Ñ%7¸×9SÑ9SÓTˆÔÜŸ™ 6×#8Ñ#8¸&×:TÑ:TÓUˆŒð 	�‰Õr7   r¬  c                 ó&   — || j                   _        y)zb
        Set the Gumbel softmax temperature to a given value. Only necessary for training
        N)r‚  r¬  )r�   r¬  s     r8   Úset_gumbel_temperaturez-Wav2Vec2ForPreTraining.set_gumbel_temperatureD  s   € ð &1ˆ�‰Õ"r7   c                 óX   — t        j                  dt        «       | j                  «        yr\  r^  r`  s    r8   ra  z/Wav2Vec2ForPreTraining.freeze_feature_extractorJ  rb  r7   c                 óL   — | j                   j                  j                  «        yrd  ©rä  rR  rß   r`  s    r8   r_  z-Wav2Vec2ForPreTraining.freeze_feature_encoderV  ó   € ð
 	�‰×'Ñ'×:Ñ:Õ<r7   Útarget_featuresÚnegative_featuresÚpredicted_featuresc                 óÈ   — t        j                  | |gd¬«      } t        j                  |j                  «       | j                  «       d¬«      j	                  | «      }||z  }|S )zé
        Compute logits for contrastive loss based using cosine similarity as the distance measure between
        `[positive_feature, negative_features]` and `[predicted_features]`. Additionally, temperature can be applied.
        r   r  rI   )r3   r  Úcosine_similarityr+  r¾  )rŠ  r‹  rŒ  r¬  Úlogitss        r8   Úcompute_contrastive_logitsz1Wav2Vec2ForPreTraining.compute_contrastive_logits]  sa   € ô  Ÿ)™) _Ð6GÐ$HÈaÔPˆä×(Ñ(Ð);×)AÑ)AÓ)CÀ_×EZÑEZÓE\ÐbdÔe×mÑmØó
ˆð
 ˜+Ñ%ˆØˆr7   )rq  rH  rä   r<   ro   ru   r  r~  r  r>   c           
      ó  — |�|n| j                   j                  }|�|j                  t        j                  «      }| j                  ||||||¬«      }| j                  |d   «      }	| j                  |d   «      }
|�!| j                  |
j                  d   |d¬«      }| j                  |
|¬«      \  }}|j                  | j                  j                  j                  «      }| j                  |«      }dx}x}}|��Ã|j                  \  }}}|j                  d|«      |j                  «       j                  d«         }|j                  ||d|«      j!                  d	ddd
«      }| j#                  |ddd…f   ||	| j                   j$                  «      }||k(  j'                  d«      }|j)                  «       rt+        d«      |dd |<   |j-                  dd	«      j/                  d|j1                  d«      «      }d|j                  «       z
  dz  j-                  dd«      j3                  «       }t4        j6                  j9                  |j+                  «       |d¬«      }| j                   j:                  | j                   j<                  z  }||z
  |z  |j?                  «       z  }|| j                   j@                  |z  z   }|s|�||	||f|d	d z   S |	||f|d	d z   S tC        ||	|||jD                  |jF                  ||¬«      S )a·  
        mask_time_indices (`torch.BoolTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Indices to mask extracted features for contrastive loss. When in training mode, model learns to predict
            masked extracted features in *config.proj_codevector_dim* space.
        sampled_negative_indices (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_negatives)`, *optional*):
            Indices indicating which quantized target vectors are used as negative sampled vectors in contrastive loss.
            Required input for pre-training.

        Returns:

        Example:

        ```python
        >>> import torch
        >>> from transformers import AutoFeatureExtractor, Wav2Vec2ForPreTraining
        >>> from transformers.models.wav2vec2.modeling_wav2vec2 import _compute_mask_indices, _sample_negative_indices
        >>> from datasets import load_dataset

        >>> feature_extractor = AutoFeatureExtractor.from_pretrained("facebook/wav2vec2-base")
        >>> model = Wav2Vec2ForPreTraining.from_pretrained("facebook/wav2vec2-base")

        >>> ds = load_dataset("hf-internal-testing/librispeech_asr_dummy", "clean", split="validation")
        >>> input_values = feature_extractor(ds[0]["audio"]["array"], return_tensors="pt").input_values  # Batch size 1

        >>> # compute masked indices
        >>> batch_size, raw_sequence_length = input_values.shape
        >>> sequence_length = model._get_feat_extract_output_lengths(raw_sequence_length).item()
        >>> mask_time_indices = _compute_mask_indices(
        ...     shape=(batch_size, sequence_length), mask_prob=0.2, mask_length=2
        ... )
        >>> sampled_negative_indices = _sample_negative_indices(
        ...     features_shape=(batch_size, sequence_length),
        ...     num_negatives=model.config.num_negatives,
        ...     mask_time_indices=mask_time_indices,
        ... )
        >>> mask_time_indices = torch.tensor(data=mask_time_indices, device=input_values.device, dtype=torch.long)
        >>> sampled_negative_indices = torch.tensor(
        ...     data=sampled_negative_indices, device=input_values.device, dtype=torch.long
        ... )

        >>> with torch.no_grad():
        ...     outputs = model(input_values, mask_time_indices=mask_time_indices)

        >>> # compute cosine similarity between predicted (=projected_states) and target (=projected_quantized_states)
        >>> cosine_sim = torch.cosine_similarity(outputs.projected_states, outputs.projected_quantized_states, dim=-1)

        >>> # show that cosine similarity is much higher than random
        >>> cosine_sim[mask_time_indices.to(torch.bool)].mean() > 0.5
        tensor(True)

        >>> # for contrastive loss training model should be put into train mode
        >>> model = model.train()
        >>> loss = model(
        ...     input_values, mask_time_indices=mask_time_indices, sampled_negative_indices=sampled_negative_indices
        ... ).loss
        ```N)r<   r  r~  ro   r  r   r    Fr
  )ro   rI   r$   r	   z-infiœÿÿÿrS   )Ú	reduction)r'   r(   r)   r*   r+   r,   r-   r.   )$r‘   rw  r>  r3   rW   rä  rê  r�  r  r9   r‚  rì  r´   rK   r  r  Úpermuter�  Úcontrastive_logits_temperatureÚallÚanyr+  r¤   ra   rq   r®  r   r  Úcross_entropyr¦  r¥  rS   Údiversity_loss_weightr&   r+   r,   )r�   rä   r<   ro   ru   r  r~  r  rk  Útransformer_featuresrv  Úquantized_featuresr*   r'   r-   r.   rc   rG   r¸   Únegative_quantized_featuresr�  Ú
neg_is_posÚtargetÚnum_codevectorss                           r8   r—   zWav2Vec2ForPreTraining.forwardr  sI  € ðJ &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ(Ø 1× 4Ñ 4´U·Z±ZÓ @Ðà—-‘-ØØ)Ø/Ø!5Ø/Ø#ð  ó 
ˆð  $×/Ñ/°¸±
Ó;Ðð  ×0Ñ0°¸±Ó<ÐàÐ%à!×DÑDØ ×&Ñ& qÑ)¨>Àuð Eó ˆNð 59·N±NØÐ0Að 5Có 5
Ñ1ÐÐ1ð 0×2Ñ2°4·>±>×3HÑ3H×3NÑ3NÓOÐØ!Ÿ^™^Ð,>Ó?Ðà37Ð7ˆÐ7Ð .Ø#Ñ/Ø7I×7OÑ7OÑ4ˆJ˜¨ð +=×*AÑ*AÀ"ÀkÓ*RØ(×-Ñ-Ó/×4Ñ4°RÓ8ñ+Ð'ð +F×*JÑ*JØ˜O¨R°ó+ç‰g�a˜˜A˜qÓ!ð (ð ×4Ñ4Ø" 4ª 7Ñ+Ø+Ø$Ø—‘×:Ñ:ó	ˆFð -Ð0KÑK×PÑPÐQSÓTˆJà�~‰~ÔÜ).¨v«��q�r�
˜:Ñ&ð ×%Ñ% a¨Ó+×3Ñ3°B¸¿¹ÀA»ÓGˆFØÐ,×1Ñ1Ó3Ñ3°tÑ;×FÑFÀqÈ!ÓL×TÑTÓVˆFä!Ÿ}™}×:Ñ:¸6¿<¹<»>È6Ð]bÐ:ÓcÐà"Ÿk™k×CÑCÀdÇkÁk×FgÑFgÑgˆOØ.Ð1FÑFÈ/ÑYÐ]n×]rÑ]rÓ]tÑtˆNð $ d§k¡k×&GÑ&GÈ.Ñ&XÑXˆDáØÐØÐ2Ð4FÐH]Ð^ÐahÐijÐikÐalÑlÐlØ(Ð*<Ð>SÐTÐW^Ð_`Ð_aÐWbÑbÐbä+ØØ1Ø'9Ø"7Ø!×/Ñ/Ø×)Ñ)Ø-Ø)ô	
ð 		
r7   )gš™™™™™¹?)NNNNNN)r/   r0   r1   r!   r„   rB   r…  ra  r_  rÆ  r3   r4   r�  r   rz  r   r&   r|  r   r,  Ú
BoolTensorrW   r   r   r—   rš   r›   s   @r8   ré  ré  5  sL  ø„ ð˜~õ ð1°#ó 1ò
&ò=ð ð
 ñ	Ø×*Ñ*ðà ×,Ñ,ðð "×-Ñ-ðð ò	ó ðñ( +Ð+DÓEÙÐ+GÐVeÔfð 26Ø8<Ø?CØ,0Ø/3Ø&*ñ^
à˜uŸ|™|Ñ,ð^
ð ! §¡Ñ.ð^
ð $ E×$4Ñ$4Ñ5ð	^
ð
 #+¨5×+;Ñ+;Ñ"<ð^
ð $ D™>ð^
ð ' t™nð^
ð ˜d‘^ð^
ð 
ˆuÐ2Ð2Ñ	3ò^
ó gó Fô^
r7   ré  z6Wav2Vec2 Model with a `language modeling` head on top.c                   óÈ   ‡ — e Zd Zˆ fd„Z ee«      	 	 	 	 	 d
dej                  deej                     dee
   dee
   dee
   deej                     deeef   fd	„«       Zˆ xZS )ÚWav2Vec2ForMaskedLMc                 ó>  •— t         ‰| �  |«       t        j                  dt        «       t        |«      | _        t        j                  |j                  «      | _
        t        j                  |j                  |j                  «      | _        | j                  «        y )NzSThe class `Wav2Vec2ForMaskedLM` is deprecated. Please use `Wav2Vec2ForCTC` instead.)rƒ   r„   ré   rê   rì   rP  rä  r   rö   Úfinal_dropoutrø   rô   r¸   r<  r  rZ  rí   s     €r8   r„   zWav2Vec2ForMaskedLM.__init__  sp   ø€ Ü‰Ñ˜Ô ä�‰ØaÔcpô	
ô & fÓ-ˆŒÜ—z‘z &×"6Ñ"6Ó7ˆŒÜ—y‘y ×!3Ñ!3°V×5FÑ5FÓGˆŒð 	�‰Õr7   rä   r<   r  r~  r  Úlabelsr>   c                 ó  — |�|n| j                   j                  }| j                  ||||¬«      }|d   }| j                  |«      }| j	                  |«      }	|s|	f|dd  z   }
|
S t        |	|j                  |j                  ¬«      S )N)r  r~  r  r   r$   )r�  r+   r,   )r‘   rw  rä  rø   r  r   r+   r,   )r�   rä   r<   r  r~  r  r¤  rk  r+   r�  Úoutputs              r8   r—   zWav2Vec2ForMaskedLM.forward%  s–   € ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—-‘-ØØ/Ø!5Ø#ð	  ó 
ˆð   ™
ˆØŸ™ ]Ó3ˆØ—‘˜mÓ,ˆáØ�Y ¨¨ Ñ,ˆFØˆMä V¸7×;PÑ;PÐ]d×]oÑ]oÔpÐpr7   ry  )r/   r0   r1   r„   r   rz  r3   r4   r   rN  rW   r,  r   r   r   r—   rš   r›   s   @r8   r¡  r¡    s±   ø„ ôñ +Ð+DÓEð 6:Ø,0Ø/3Ø&*Ø)-ñqà×'Ñ'ðqð ! ×!1Ñ!1Ñ2ðqð $ D™>ð	qð
 ' t™nðqð ˜d‘^ðqð ˜Ÿ™Ñ&ðqð 
ˆu�nÐ$Ñ	%òqó Fôqr7   r¡  zfWav2Vec2 Model with a `language modeling` head on top for Connectionist Temporal Classification (CTC).a.  
        target_lang (`str`, *optional*):
            Language id of adapter weights. Adapter weights are stored in the format adapter.<lang>.safetensors or
            adapter.<lang>.bin. Only relevant when using an instance of [`Wav2Vec2ForCTC`] with adapters. Uses 'eng' by
            default.
    c                   ó  ‡ — e Zd Zddee   fˆ fd„Zd„ Zd„ Zd„ Zd„ Z	 e
e«       eeeeee¬«      	 	 	 	 	 ddeej&                     d	eej&                     d
ee   dee   dee   deej&                     deeef   fd„«       «       Zˆ xZS )r  r  c                 ó®  •— t         ‰| �  |«       t        |«      | _        t	        j
                  |j                  «      | _        || _        |j                  €t        d| j                  › d�«      ‚t        |d«      r|j                  r|j                  n|j                  }t	        j                   ||j                  «      | _        | j%                  «        y )NzYou are trying to instantiate z÷ with a configuration that does not define the vocabulary size of the language model head. Please instantiate the model as follows: `Wav2Vec2ForCTC.from_pretrained(..., vocab_size=vocab_size)`. or define `vocab_size` of your model's configuration.rþ  )rƒ   r„   rP  rä  r   rö   r£  rø   r  r<  rM   r“   r¼   rþ  rÌ  r¸   rô   r  rZ  )r�   r‘   r  rÌ  r“   s       €r8   r„   zWav2Vec2ForCTC.__init__N  s½   ø€ Ü‰Ñ˜Ô ä% fÓ-ˆŒÜ—z‘z &×"6Ñ"6Ó7ˆŒà&ˆÔà×ÑÐ$ÜØ0°·±Ð0@ð AHð Hóð ô *1°¸Ô)GÈF×L^ÒL^ˆF×%Ò%Ðdj×dvÑdvð 	ô —y‘yÐ!3°V×5FÑ5FÓGˆŒð 	�‰Õr7   c                 óö   — | j                   }|�&t        | j                  dd«      €t        d|› d�«      ‚|€-t        | j                  dd«      �t        j                  d«       y|�| j                  |d¬«       yy)a'  
        This method overwrites [`~PreTrainedModel.tie_weights`] so that adapter weights can be correctly loaded when
        passing `target_lang=...` to `from_pretrained(...)`.

        This method is **not** supposed to be called by the user and is prone to be changed in the future.
        Nro  zCannot pass `target_lang`: r   z)By default `target_lang` is set to 'eng'.T)r?  )r  rp  r‘   rM   r<  ÚinforG  )r�   r  s     r8   Útie_weightszWav2Vec2ForCTC.tie_weightse  sƒ   € ð ×&Ñ&ˆàÐ"¤w¨t¯{©{Ð<NÐPTÓ'UÐ']ÜÐ:¸;¸-ÐGtÐuÓvÐvØÐ ¤W¨T¯[©[Ð:LÈdÓ%SÐ%_Ü�K‰KÐCÕDØÐ$Ø×Ñ˜k°dÐÕ;ð %r7   c                 óX   — t        j                  dt        «       | j                  «        y©re  r]  Nr^  r`  s    r8   ra  z'Wav2Vec2ForCTC.freeze_feature_extractorz  rb  r7   c                 óL   — | j                   j                  j                  «        yrd  rˆ  r`  s    r8   r_  z%Wav2Vec2ForCTC.freeze_feature_encoder†  r‰  r7   c                 óP   — | j                   j                  «       D ]	  }d|_        Œ y©zÒ
        Calling this function will disable the gradient computation for the base model so that its parameters will not
        be updated during training. Only the classification head will be updated.
        FN©rä  rÛ   rÜ   rÝ   s     r8   Úfreeze_base_modelz Wav2Vec2ForCTC.freeze_base_model�  ó(   € ð
 —]‘]×-Ñ-Ó/ò 	(ˆEØ"'ˆEÕñ	(r7   )rp  rq  rH  rs  Úexpected_lossrä   r<   r  r~  r  r¤  r>   c           
      ó¤  — |�|n| j                   j                  }|�I|j                  «       | j                   j                  k\  r"t	        d| j                   j                  › �«      ‚| j                  |||||¬«      }|d   }| j                  |«      }| j                  |«      }	d}
|��b|�|n$t        j                  |t        j                  ¬«      }| j                  |j                  d«      «      j                  t        j                  «      }|dk\  }|j                  d«      }|j                  |«      }t        j                   j#                  |	dt        j$                  ¬«      j'                  dd«      }t        j(                  j*                  j-                  d	¬
«      5  t        j                   j/                  ||||| j                   j0                  | j                   j2                  | j                   j4                  ¬«      }
ddd«       |s|	f|t6        d z   }|
�|
f|z   S |S t9        |
|	|j:                  |j<                  ¬«      S # 1 sw Y   ŒExY w)aà  
        labels (`torch.LongTensor` of shape `(batch_size, target_length)`, *optional*):
            Labels for connectionist temporal classification. Note that `target_length` has to be smaller or equal to
            the sequence length of the output logits. Indices are selected in `[-100, 0, ..., config.vocab_size - 1]`.
            All labels set to `-100` are ignored (masked), the loss is only computed for labels in `[0, ...,
            config.vocab_size - 1]`.
        Nz$Label values must be <= vocab_size: ru  r   rJ   rI   )r¶   rK   r    F)Úenabled)Úblankr’  Úzero_infinity©r'   r�  r+   r,   )r‘   rw  rC   r<  rM   rä  rø   r  r3   Ú	ones_liker  r  rS   r>  Úmasked_selectr   r  Úlog_softmaxr9  r¤   ÚbackendsÚcudnnÚflagsÚctc_lossÚpad_token_idÚctc_loss_reductionÚctc_zero_infinityÚ_HIDDEN_STATES_START_POSITIONr   r+   r,   )r�   rä   r<   r  r~  r  r¤  rk  r+   r�  r'   re   Úlabels_maskÚtarget_lengthsÚflattened_targetsÚ	log_probsr¦  s                    r8   r—   zWav2Vec2ForCTC.forward•  s'  € ð0 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ &§*¡*£,°$·+±+×2HÑ2HÒ"HÜÐCÀDÇKÁK×DZÑDZÐC[Ð\Ó]Ð]à—-‘-ØØ)Ø/Ø!5Ø#ð  ó 
ˆð   ™
ˆØŸ™ ]Ó3ˆà—‘˜mÓ,ˆàˆØÑð #1Ð"<‘Ä%Ç/Á/ÐR^Ôfk×fpÑfpÔBqð ð !×AÑAÀ.×BTÑBTÐUWÓBXÓY×\Ñ\Ô]b×]gÑ]gÓhˆMð ! A™+ˆKØ(Ÿ_™_¨RÓ0ˆNØ &× 4Ñ 4°[Ó AÐô Ÿ™×1Ñ1°&¸bÌÏÉÐ1ÓV×`Ñ`ÐabÐdeÓfˆIä—‘×%Ñ%×+Ñ+°EÐ+Ó:ñ 	Ü—}‘}×-Ñ-ØØ%Ø!Ø"ØŸ+™+×2Ñ2Ø"Ÿk™k×<Ñ<Ø"&§+¡+×"?Ñ"?ð .ó �÷	ñ Ø�Y Ô)FÐ)GÐ!HÑHˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEäØ˜f°G×4IÑ4IÐV]×VhÑVhô
ð 	
÷	ð 	ús   ÆA#IÉIr•   ry  )r/   r0   r1   r   rU  r„   r«  ra  r_  r²  r   rz  r   r{  r   r|  Ú_CTC_EXPECTED_OUTPUTÚ_CTC_EXPECTED_LOSSr3   r,  rW   r   r   r—   rš   r›   s   @r8   r  r  C  sí   ø„ ñ¨H°S©Mõ ò.<ò*
&ò=ò(ñ +Ð+DÓEÙØ&Ø"Ø$Ø,Ø(ôð 26Ø,0Ø/3Ø&*Ø)-ñD
à˜uŸ|™|Ñ,ðD
ð ! §¡Ñ.ðD
ð $ D™>ð	D
ð
 ' t™nðD
ð ˜d‘^ðD
ð ˜Ÿ™Ñ&ðD
ð 
ˆu�nÐ$Ñ	%òD
óó FôD
r7   r  z—
    Wav2Vec2 Model with a sequence classification head on top (a linear layer over the pooled output) for tasks like
    SUPERB Keyword Spotting.
    c                   ó  ‡ — e Zd Zˆ fd„Zd„ Zd„ Zd„ Z ee«       e	e
eedee¬«      	 	 	 	 	 ddeej"                     deej"                     d	ee   d
ee   dee   deej"                     deeef   fd„«       «       Zˆ xZS )Ú!Wav2Vec2ForSequenceClassificationc                 óü  •— t         ‰| �  |«       t        |d«      r|j                  rt	        d«      ‚t        |«      | _        |j                  dz   }|j                  r0t        j                  t        j                  |«      |z  «      | _        t        j                  |j                  |j                   «      | _        t        j                  |j                   |j$                  «      | _        | j)                  «        y )Nrþ  z_Sequence classification does not support the use of Wav2Vec2 adapters (config.add_adapter=True)r    )rƒ   r„   r¼   rþ  rM   rP  rä  ry  Úuse_weighted_layer_sumr   r©  r3   r\   Úlayer_weightsrô   r¸   Úclassifier_proj_sizeÚ	projectorÚ
num_labelsÚ
classifierrZ  ©r�   r‘   Ú
num_layersr“   s      €r8   r„   z*Wav2Vec2ForSequenceClassification.__init__ì  sÀ   ø€ Ü‰Ñ˜Ô ä�6˜=Ô)¨f×.@Ò.@ÜØqóð ô & fÓ-ˆŒØ×-Ñ-°Ñ1ˆ
Ø×(Ò(Ü!#§¡¬e¯j©j¸Ó.DÀzÑ.QÓ!RˆDÔÜŸ™ 6×#5Ñ#5°v×7RÑ7RÓSˆŒÜŸ)™) F×$?Ñ$?À×ARÑARÓSˆŒð 	�‰Õr7   c                 óX   — t        j                  dt        «       | j                  «        yr\  r^  r`  s    r8   ra  z:Wav2Vec2ForSequenceClassification.freeze_feature_extractorý  rb  r7   c                 óL   — | j                   j                  j                  «        yrd  rˆ  r`  s    r8   r_  z8Wav2Vec2ForSequenceClassification.freeze_feature_encoder		  r‰  r7   c                 óP   — | j                   j                  «       D ]	  }d|_        Œ yr°  r±  rÝ   s     r8   r²  z3Wav2Vec2ForSequenceClassification.freeze_base_model	  r³  r7   rn  )rp  rq  rH  rr  rs  r´  rä   r<   r  r~  r  r¤  r>   c                 ó<  — |�|n| j                   j                  }| j                   j                  rdn|}| j                  |||||¬«      }| j                   j                  rr|t           }t        j                  |d¬«      }t        j                  j                  | j                  d¬«      }	||	j                  ddd«      z  j                  d¬«      }n|d   }| j                  |«      }|€|j                  d¬«      }
n‰| j                  |j                   d   |«      }|j#                  d«      j%                  dd|j                   d   «      }d	|| <   |j                  d¬«      |j                  d¬«      j                  dd«      z  }
| j'                  |
«      }d}|�Ft)        «       } ||j                  d| j                   j*                  «      |j                  d«      «      }|s|f|t        d z   }|�|f|z   S |S t-        |||j.                  |j0                  ¬
«      S )á�  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
        NTru  r    r  rI   r   r$   r)  r¹  )r‘   rw  rÎ  rä  rÄ  r3   Ústackr   r  r  rÏ  r  rS   rÑ  r±  r  r9   rŒ  r�  rÓ  r   rÒ  r   r+   r,   )r�   rä   r<   r  r~  r  r¤  rk  r+   Únorm_weightsÚpooled_outputÚpadding_maskÚexpand_padding_maskr�  r'   Úloss_fctr¦  s                    r8   r—   z)Wav2Vec2ForSequenceClassification.forward	  s  € ð2 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ'+§{¡{×'IÒ'I™tÐOcÐà—-‘-ØØ)Ø/Ø!5Ø#ð  ó 
ˆð �;‰;×-Ò-Ø#Ô$AÑBˆMÜ!ŸK™K¨¸1Ô=ˆMÜŸ=™=×0Ñ0°×1CÑ1CÈÐ0ÓLˆLØ*¨\×->Ñ->¸rÀ1ÀaÓ-HÑH×MÑMÐRSÐMÓT‰Mà# A™JˆMàŸ™ }Ó5ˆØÐ!Ø)×.Ñ.°1Ð.Ó5‰Mà×BÑBÀ=×CVÑCVÐWXÑCYÐ[iÓjˆLØ".×"8Ñ"8¸Ó"<×"CÑ"CÀAÀqÈ-×J]ÑJ]Ð^_ÑJ`Ó"aÐØ25ˆMÐ.Ð.Ñ/Ø)×-Ñ-°!Ð-Ó4°|×7GÑ7GÈAÐ7GÓ7N×7SÑ7SÐTVÐXYÓ7ZÑZˆMà—‘ Ó/ˆàˆØÐÜ'Ó)ˆHÙ˜FŸK™K¨¨D¯K©K×,BÑ,BÓCÀVÇ[Á[ÐQSÃ_ÓUˆDáØ�Y Ô)FÐ)GÐ!HÑHˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä'ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r7   ry  )r/   r0   r1   r„   ra  r_  r²  r   rz  r   Ú_SEQ_CLASS_CHECKPOINTr   r|  Ú_SEQ_CLASS_EXPECTED_OUTPUTÚ_SEQ_CLASS_EXPECTED_LOSSr   r3   r,  rW   r   r   r—   rš   r›   s   @r8   rÌ  rÌ  ä  sØ   ø„ ôò"
&ò=ò(ñ +Ð+DÓEÙØ(Ø,Ø$ØØ2Ø.ôð 26Ø,0Ø/3Ø&*Ø)-ñ<
à˜uŸ|™|Ñ,ð<
ð ! §¡Ñ.ð<
ð $ D™>ð	<
ð
 ' t™nð<
ð ˜d‘^ð<
ð ˜Ÿ™Ñ&ð<
ð 
ˆuÐ.Ð.Ñ	/ò<
óó Fô<
r7   rÌ  zd
    Wav2Vec2 Model with a frame classification head on top for tasks like Speaker Diarization.
    c                   ó   ‡ — e Zd Zˆ fd„Zd„ Zd„ Zd„ Z ee«       e	e
eede¬«      	 	 	 	 	 ddeej                      deej                      d	eej                      d
ee   dee   dee   deeef   fd„«       «       Zˆ xZS )Ú#Wav2Vec2ForAudioFrameClassificationc                 óÀ  •— t         ‰| �  |«       t        |d«      r|j                  rt	        d«      ‚t        |«      | _        |j                  dz   }|j                  r0t        j                  t        j                  |«      |z  «      | _        t        j                  |j                  |j                   «      | _        |j                   | _        | j%                  «        y )Nrþ  zbAudio frame classification does not support the use of Wav2Vec2 adapters (config.add_adapter=True)r    )rƒ   r„   r¼   rþ  rM   rP  rä  ry  rÎ  r   r©  r3   r\   rÏ  rô   r¸   rÒ  rÓ  Úinit_weightsrÔ  s      €r8   r„   z,Wav2Vec2ForAudioFrameClassification.__init__g	  s¯   ø€ Ü‰Ñ˜Ô ä�6˜=Ô)¨f×.@Ò.@ÜØtóð ô & fÓ-ˆŒØ×-Ñ-°Ñ1ˆ
Ø×(Ò(Ü!#§¡¬e¯j©j¸Ó.DÀzÑ.QÓ!RˆDÔÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒØ ×+Ñ+ˆŒà×ÑÕr7   c                 óX   — t        j                  dt        «       | j                  «        yr­  r^  r`  s    r8   ra  z<Wav2Vec2ForAudioFrameClassification.freeze_feature_extractorw	  rb  r7   c                 óL   — | j                   j                  j                  «        yrd  rˆ  r`  s    r8   r_  z:Wav2Vec2ForAudioFrameClassification.freeze_feature_encoderƒ	  r‰  r7   c                 óP   — | j                   j                  «       D ]	  }d|_        Œ yr°  r±  rÝ   s     r8   r²  z5Wav2Vec2ForAudioFrameClassification.freeze_base_modelŠ	  r³  r7   rn  ro  rä   r<   r¤  r  r~  r  r>   c           	      óú  — |�|n| j                   j                  }| j                   j                  rdn|}| j                  |||||¬«      }| j                   j                  rr|t           }t        j                  |d¬«      }t        j                  j                  | j                  d¬«      }	||	j                  ddd«      z  j                  d¬«      }n|d   }| j                  |«      }
d}|�\t        «       } ||
j                  d| j                  «      t        j                   |j                  d| j                  «      d¬«      «      }|s|
f|t        d z   }|S t#        ||
|j$                  |j&                  ¬	«      S )
rÚ  NTru  r    r  rI   r   )Úaxisr¹  )r‘   rw  rÎ  rä  rÄ  r3   rÛ  r   r  r  rÏ  r  rS   rÓ  r   rÒ  r¿  r   r+   r,   )r�   rä   r<   r¤  r  r~  r  rk  r+   rÜ  r�  r'   rà  r¦  s                 r8   r—   z+Wav2Vec2ForAudioFrameClassification.forward’	  sh  € ð0 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ'+§{¡{×'IÒ'I™tÐOcÐà—-‘-ØØ)Ø/Ø!5Ø#ð  ó 
ˆð �;‰;×-Ò-Ø#Ô$AÑBˆMÜ!ŸK™K¨¸1Ô=ˆMÜŸ=™=×0Ñ0°×1CÑ1CÈÐ0ÓLˆLØ*¨\×->Ñ->¸rÀ1ÀaÓ-HÑH×MÑMÐRSÐMÓT‰Mà# A™JˆMà—‘ Ó/ˆàˆØÐÜ'Ó)ˆHÙ˜FŸK™K¨¨D¯O©OÓ<¼e¿l¹lÈ6Ï;É;ÐWYÐ[_×[jÑ[jÓKkÐrsÔ>tÓuˆDáØ�Y Ô)FÐ)GÐ!HÑHˆFØˆMä$ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r7   ry  )r/   r0   r1   r„   ra  r_  r²  r   rz  r   Ú_FRAME_CLASS_CHECKPOINTr   r|  Ú_FRAME_EXPECTED_OUTPUTr   r3   r,  rW   r   r   r—   rš   r›   s   @r8   rå  rå  `	  sÕ   ø„ ôò 
&ò=ò(ñ +Ð+DÓEÙØ*Ø)Ø$ØØ.ôð 26Ø)-Ø,0Ø/3Ø&*ñ3
à˜uŸ|™|Ñ,ð3
ð ! §¡Ñ.ð3
ð ˜Ÿ™Ñ&ð	3
ð
 $ D™>ð3
ð ' t™nð3
ð ˜d‘^ð3
ð 
ˆuÐ+Ð+Ñ	,ò3
óó Fô3
r7   rå  c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )ÚAMSoftmaxLossc                 óæ   •— t         t        | �  «        || _        || _        || _        t        j                  t        j                  ||«      d¬«      | _
        t        j                  «       | _        y )NT)rÜ   )rƒ   rð  r„   ÚscaleÚmarginrÒ  r   r©  r3   Úrandnr´   r   r'   )r�   rÚ  rÒ  rò  ró  r“   s        €r8   r„   zAMSoftmaxLoss.__init__Ñ	  sS   ø€ ÜŒm˜TÑ+Ô-ØˆŒ
ØˆŒØ$ˆŒÜ—l‘l¤5§;¡;¨y¸*Ó#EÐUYÔZˆŒÜ×'Ñ'Ó)ˆ�	r7   c                 óä  — |j                  «       }t        j                  j                  | j                  d¬«      }t        j                  j                  |d¬«      }t        j                  ||«      }|| j                  z
  }t        j                  j                  || j                  «      }| j                  t        j                  |j                  «       ||«      z  }| j                  ||«      }|S )Nr   r  r    )r®  r   r  Ú	normalizer´   r3   Úmmró  Úone_hotrÒ  rò  r¯  rW   r'   )	r�   r+   r¤  r´   Ú	cos_thetaÚpsiÚonehotr�  r'   s	            r8   r—   zAMSoftmaxLoss.forwardÙ	  s²   € Ø—‘Ó!ˆÜ—‘×(Ñ(¨¯©¸!Ð(Ó<ˆÜŸ™×/Ñ/°À1Ð/ÓEˆÜ—H‘H˜]¨FÓ3ˆ	Ø˜$Ÿ+™+Ñ%ˆä—‘×&Ñ& v¨t¯©Ó?ˆØ—‘œeŸk™k¨&¯+©+«-¸¸iÓHÑHˆØ�y‰y˜ Ó(ˆàˆr7   )g      >@gš™™™™™Ù?r™   r›   s   @r8   rð  rð  Ð	  s   ø„ õ*ör7   rð  c                   óX   ‡ — e Zd Zdˆ fd„	Zdej
                  dej
                  fd„Zˆ xZS )Ú	TDNNLayerc                 óš  •— t         ‰| �  «        |dkD  r|j                  |dz
     n|j                  |   | _        |j                  |   | _        |j
                  |   | _        |j                  |   | _        t        j                  | j                  | j                  z  | j                  «      | _        t        j                  «       | _        y )Nr   r    )rƒ   r„   Útdnn_dimr†   r‡   Útdnn_kernelr€   Útdnn_dilationÚdilationr   rô   ÚkernelrÞ  rŽ   r�   s      €r8   r„   zTDNNLayer.__init__è	  s¡   ø€ Ü‰ÑÔØ<DÀqºL˜6Ÿ?™?¨8°a©<Ò8ÈfÏoÉoÐ^fÑNgˆÔØ"ŸO™O¨HÑ5ˆÔØ!×-Ñ-¨hÑ7ˆÔØ×,Ñ,¨XÑ6ˆŒä—i‘i × 0Ñ 0°4×3CÑ3CÑ CÀT×EVÑEVÓWˆŒÜŸ'™'›)ˆ�r7   r+   r>   c                 ó&  — t        «       rddlm} t        «       r+t        | j                  «      rt        j                  d«       |j                  dd«      }| j                  j                  j                  | j                  | j                  | j                  «      j                  dd«      }t        j                  j                  ||| j                  j                   | j"                  ¬«      }|j                  dd«      }| j%                  |«      }|S )Nr   )Ú	LoraLayerz‡Detected LoRA on TDNNLayer. LoRA weights won't be applied due to optimization. You should exclude TDNNLayer from LoRA's target modules.r    r$   )r  )r   Úpeft.tuners.lorar  rS  r  ré   rê   r¤   r´   r  r‡   r€   r†   r   r  Úconv1dr‚   r  rŽ   )r�   r+   r  r´   s       r8   r—   zTDNNLayer.forwardò	  sØ   € ÜÔÝ2äÔÜ˜$Ÿ+™+ yÔ1Ü—‘ðOôð &×/Ñ/°°1Ó5ˆØ—‘×#Ñ#×(Ñ(¨×):Ñ):¸D×<LÑ<LÈd×N^ÑN^Ó_×iÑiÐjkÐmnÓoˆÜŸ™×,Ñ,¨]¸FÀDÇKÁK×DTÑDTÐ_c×_lÑ_lÐ,ÓmˆØ%×/Ñ/°°1Ó5ˆàŸ™¨Ó6ˆØÐr7   r˜   )r/   r0   r1   r„   r3   r,  r—   rš   r›   s   @r8   rý  rý  ç	  s#   ø„ õ$ð U§\¡\ð °e·l±l÷ r7   rý  zl
    Wav2Vec2 Model with an XVector feature extraction head on top for tasks like Speaker Verification.
    c                   ó*  ‡ — e Zd Zˆ fd„Zd„ Zd„ Zd„ Zdeej                  e
f   fd„Z ee«       eeeede¬«      	 	 	 	 	 dd	eej(                     d
eej(                     dee   dee   dee   deej(                     deeef   fd„«       «       Zˆ xZS )ÚWav2Vec2ForXVectorc                 ó  •— t         ‰| �  |«       t        |«      | _        |j                  dz   }|j
                  r0t        j                  t        j                  |«      |z  «      | _
        t        j                  |j                  |j                  d   «      | _        t        t!        |j                  «      «      D �cg c]  }t#        ||«      ‘Œ }}t        j$                  |«      | _        t        j                  |j                  d   dz  |j(                  «      | _        t        j                  |j(                  |j(                  «      | _        t/        |j(                  |j0                  «      | _        | j5                  «        y c c}w )Nr    r   rI   r$   )rƒ   r„   rP  rä  ry  rÎ  r   r©  r3   r\   rÏ  rô   r¸   rÿ  rÑ  rU   rZ   rý  rÔ   ÚtdnnÚxvector_output_dimrR  rÓ  rð  rÒ  Ú	objectiverç  )r�   r‘   rÕ  rØ   Útdnn_layersr“   s        €r8   r„   zWav2Vec2ForXVector.__init__
  s  ø€ Ü‰Ñ˜Ô ä% fÓ-ˆŒØ×-Ñ-°Ñ1ˆ
Ø×(Ò(Ü!#§¡¬e¯j©j¸Ó.DÀzÑ.QÓ!RˆDÔÜŸ™ 6×#5Ñ#5°v·±ÀqÑ7IÓJˆŒä5:¼3¸v¿¹Ó;OÓ5PÖQ°”y ¨Õ+ÐQˆÐQÜ—M‘M +Ó.ˆŒ	ä!#§¡¨6¯?©?¸2Ñ+>ÀÑ+BÀF×D]ÑD]Ó!^ˆÔÜŸ)™) F×$=Ñ$=¸v×?XÑ?XÓYˆŒä& v×'@Ñ'@À&×BSÑBSÓTˆŒà×ÑÕùò Rs   Â>Fc                 óX   — t        j                  dt        «       | j                  «        yr­  r^  r`  s    r8   ra  z+Wav2Vec2ForXVector.freeze_feature_extractor!
  rb  r7   c                 óL   — | j                   j                  j                  «        yrd  rˆ  r`  s    r8   r_  z)Wav2Vec2ForXVector.freeze_feature_encoder-
  r‰  r7   c                 óP   — | j                   j                  «       D ]	  }d|_        Œ yr°  r±  rÝ   s     r8   r²  z$Wav2Vec2ForXVector.freeze_base_model4
  r³  r7   re   c                 óV   — d„ }| j                   j                  D ]  } |||d«      }Œ |S )z?
        Computes the output length of the TDNN layers
        c                 ó   — | |z
  |z  dz   S )Nr    r6   r  s      r8   r  zEWav2Vec2ForXVector._get_tdnn_output_lengths.<locals>._conv_out_lengthA
  s   € ð ! ;Ñ.°6Ñ9¸AÑ=Ð=r7   r    )r‘   r   )r�   re   r  r€   s       r8   Ú_get_tdnn_output_lengthsz+Wav2Vec2ForXVector._get_tdnn_output_lengths<
  s:   € ò
	>ð
  Ÿ;™;×2Ñ2ò 	LˆKÙ,¨]¸KÈÓK‰Mð	Lð Ðr7   rn  ro  rä   r<   r  r~  r  r¤  r>   c                 óö  — |�|n| j                   j                  }| j                   j                  rdn|}| j                  |||||¬«      }| j                   j                  rr|t           }t        j                  |d¬«      }t        j                  j                  | j                  d¬«      }	||	j                  ddd«      z  j                  d¬«      }n|d   }| j                  |«      }| j                  D ]
  }
 |
|«      }Œ |€%|j                  d¬«      }|j!                  d¬«      }nÃ| j#                  |j                  d¬«      «      }| j%                  |«      }g }g }t'        |«      D ]U  \  }}|j)                  ||d|…f   j                  d¬«      «       |j)                  ||d|…f   j!                  d¬«      «       ŒW t        j                  |«      }t        j                  |«      }t        j*                  ||gd¬«      }| j-                  |«      }| j/                  |«      }d}|�| j1                  ||«      }|s||f|t        d z   }|�|f|z   S |S t3        ||||j4                  |j6                  ¬«      S )	rÚ  NTru  r    r  rI   r   )r'   r�  Ú
embeddingsr+   r,   )r‘   rw  rÎ  rä  rÄ  r3   rÛ  r   r  r  rÏ  r  rS   rÑ  r  r±  ræ  r  r  Ú	enumerater^   r  rR  rÓ  r  r   r+   r,   )r�   rä   r<   r  r~  r  r¤  rk  r+   rÜ  Ú
tdnn_layerÚmean_featuresÚstd_featuresÚfeat_extract_output_lengthsÚtdnn_output_lengthsrØ   ÚlengthÚstatistic_poolingÚoutput_embeddingsr�  r'   r¦  s                         r8   r—   zWav2Vec2ForXVector.forwardK
  s™  € ð0 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ'+§{¡{×'IÒ'I™tÐOcÐà—-‘-ØØ)Ø/Ø!5Ø#ð  ó 
ˆð �;‰;×-Ò-Ø#Ô$AÑBˆMÜ!ŸK™K¨¸1Ô=ˆMÜŸ=™=×0Ñ0°×1CÑ1CÈÐ0ÓLˆLØ*¨\×->Ñ->¸rÀ1ÀaÓ-HÑH×MÑMÐRSÐMÓT‰Mà# A™JˆMàŸ™ }Ó5ˆàŸ)™)ò 	6ˆJÙ& }Ó5‰Mð	6ð Ð!Ø)×.Ñ.°1Ð.Ó5ˆMØ(×,Ñ,°Ð,Ó3‰Là*.×*OÑ*OÐP^×PbÑPbÐghÐPbÓPiÓ*jÐ'Ø"&×"?Ñ"?Ð@[Ó"\ÐØˆMØˆLÜ&Ð':Ó;ò J‘	��6Ø×$Ñ$ ]°1°g°v°g°:Ñ%>×%CÑ%CÈÐ%CÓ%JÔKØ×#Ñ# M°!°W°f°W°*Ñ$=×$AÑ$AÀaÐ$AÓ$HÕIðJô "ŸK™K¨Ó6ˆMÜ Ÿ;™; |Ó4ˆLÜ!ŸI™I }°lÐ&CÈÔLÐà ×2Ñ2Ð3DÓEÐØ—‘Ð!2Ó3ˆàˆØÐØ—>‘> &¨&Ó1ˆDáØÐ/Ð0°7Ô;XÐ;YÐ3ZÑZˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEäØØØ(Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r7   ry  )r/   r0   r1   r„   ra  r_  r²  r   r3   rN  rB   r  r   rz  r   Ú_XVECTOR_CHECKPOINTr   r|  Ú_XVECTOR_EXPECTED_OUTPUTr   r,  rW   r   r—   rš   r›   s   @r8   r	  r	  
  sù   ø„ ôò&
&ò=ò(ð°e¸E×<LÑ<LÈcÐ<QÑ6Ró ñ +Ð+DÓEÙØ&Ø!Ø$ØØ0ôð 26Ø,0Ø/3Ø&*Ø)-ñI
à˜uŸ|™|Ñ,ðI
ð ! §¡Ñ.ðI
ð $ D™>ð	I
ð
 ' t™nðI
ð ˜d‘^ðI
ð ˜Ÿ™Ñ&ðI
ð 
ˆu�mÐ#Ñ	$òI
óó FôI
r7   r	  )rå  r  r¡  ré  rÌ  r	  rP  rã  rÌ   r•   )or2   ró  ré   Údataclassesr   Útypingr   r   r   ÚnumpyrN   r3   Útorch.utils.checkpointr   Útorch.nnr   Úactivationsr
   Úintegrations.deepspeedr   Úintegrations.fsdpr   Úmodeling_flash_attention_utilsr   r   Úmodeling_outputsr   r   r   r   r   r   r   Úmodeling_utilsr   r»   r   r   r   r   r   r   r   r   r   Úconfiguration_wav2vec2r!   r8  r3  Úsafetensors.torchr"   r5  r#   Ú
get_loggerr/   r<  rÄ  r|  r{  r}  rÉ  rÊ  rá  râ  rã  rí  rî  r   r!  r&   rB   r+  rN  Úndarrayrl   r{   ÚModuler}   r�   r¦   r®   rÅ   rÎ   rç   rï   rü   r.  rD  r_  rM  r\  rm  ru  rž  r£  rÈ  rË  rq  rã  ÚWAV2VEC2_START_DOCSTRINGrz  rP  ré  r¡  r  rÌ  rå  rð  rý  r	  Ú__all__r6   r7   r8   ú<module>r4     s†  ðñ ã Û Ý !ß )Ñ )ã Û Û Ý Ý %å !Ý @Ý 7ß h÷÷ ñ õ .÷
÷ 
õ 
õ 3ð ,Ð Ø5Ð áÔÝ=ñ ÔÝJð 
ˆ×	Ñ	˜HÓ	%€ð !"Ð ð #€ð 4Ð Ú&Ð ð uÐ ØÐ ð 9Ð Ø*Ð ØÐ ð <Ð Ø˜Q˜Ð ð 8Ð ØÐ ð ô&7 ;ó &7ó ð&7ðZ 26ØñtØ��c�‰?ðtàðtð ðtð ˜U×-Ñ-Ñ.ð	tð
 ðtð ‡Z�Zótðp Z^ñ!$Øð!$Ø*-ð!$ØBJÈ2Ï:É:ÑBVó!$ôH 2§9¡9ô ô* §¡ô ô6 §¡ô ô0* b§i¡iô *ôZ˜2Ÿ9™9ô ô+˜RŸY™Yô +ô\
Ð5ô 
ô1 §	¡	ô 1ô [B˜Ÿ	™	ô [Bô~{9Ð/ô {9ô|h1Ð-ô h1ðX Ø!Ø0ñÐ ô˜"Ÿ)™)ô ô0 ˜2Ÿ9™9ô  ôF*¨"¯)©)ô *ôZR
�b—i‘iô R
ôjV
 R§Y¡Yô V
ôrI' B§I¡Iô I'ôX�b—i‘iô ô>˜2Ÿ9™9ô ô$˜rŸy™yô ô2x'˜oô x'ðv	Ð ð&#Ð ñL ØhØóôP
Ð+ó P
ó	ðP
ñf ÐQÐSkÓlô\
Ð4ó \
ó mð\
ñ~ ÐRÐTlÓmô*qÐ1ó *qó nð*qñZ ØpØðó	ôT
Ð,ó T
ó	ðT
ñn ðð óôr
Ð(?ó r
óðr
ñj ðð ó	ôg
Ð*Aó g
óðg
ôT�B—I‘Iô ô.�—	‘	ô ñ@ ðð ó	ôO
Ð0ó O
óðO
òd	�r7   