Ë
    T^(hIà  ã                   ó\  — d Z ddlmZ ddlmZmZmZ ddlZddlm	Z
 ddlZddlmZ ddlZddlmZmZmZ ddlmZ ddlmZ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! d	dl"m#Z#m$Z$m%Z%m&Z& ddl'm(Z(  e&jR                  e*«      Z+ejX                  jZ                   G d„ de#«      «       Z.ejX                  jZ                   G d„ de#«      «       Z/	 	 dSdee0e0f   de1de0deejd                     de0dejd                  fd„Z3dTdede0deejd                     fd„Z4dZ5dZ6 G d„ d e
jn                  «      Z8 G d!„ d"e
jn                  «      Z9 G d#„ d$e
jn                  «      Z: G d%„ d&e
jn                  «      Z; G d'„ d(e
jn                  «      Z< G d)„ d*e
jn                  «      Z= G d+„ d,e
jn                  «      Z> G d-„ d.e
jn                  «      Z? G d/„ d0e
jn                  «      Z@ G d1„ d2e
jn                  «      ZA G d3„ d4e
jn                  «      ZB G d5„ d6e
jn                  «      ZC G d7„ d8e
jn                  «      ZD G d9„ d:e
jn                  «      ZE G d;„ d<e
jn                  «      ZF G d=„ d>e«      ZG G d?„ d@e
jn                  «      ZH e$dAe5«       G dB„ dCeG«      «       ZIdDZJ e!eIe6eJz   «        e eIe.e(¬E«        G dF„ dGe
jn                  «      ZK e$dHe5«       G dI„ dJeG«      «       ZLdKZM e!eLe6eMz   «        e eLee(¬E«        G dL„ dMe
jn                  «      ZN e$dNe5«       G dO„ dPeG«      «       ZOdQZP e!eOe6ePz   «        e eOe/e(¬E«       g dR¢ZQy)UzFlax Wav2Vec2 model.é    )Úpartial)ÚOptionalÚTupleÚUnionN)Ú
FrozenDictÚfreezeÚunfreeze)Údot_product_attention_weights)Úflatten_dictÚunflatten_dict)Úlaxé   )ÚFlaxBaseModelOutputÚFlaxCausalLMOutput)ÚACT2FNÚFlaxPreTrainedModelÚ append_replace_return_docstringsÚoverwrite_call_docstring)ÚModelOutputÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingé   )ÚWav2Vec2Configc                   ó²   — e Zd ZU dZdZej                  ed<   dZej                  ed<   dZ	e
eej                        ed<   dZe
eej                        ed<   y)ÚFlaxWav2Vec2BaseModelOutputa•  
    Output type of [`FlaxWav2Vec2BaseModelOutput`], with potential hidden states and attentions.

    Args:
        last_hidden_state (`jnp.ndarray` of shape `(batch_size, sequence_length, hidden_size)`):
            Sequence of hidden-states at the output of the last layer of the model.
        extract_features (`jnp.ndarray` of shape `(batch_size, sequence_length, last_conv_dim)`):
            Sequence of extracted feature vectors of the last convolutional layer of the model with `last_conv_dim`
            being the dimension of the last convolutional layer.
        hidden_states (`tuple(jnp.ndarray)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `jnp.ndarray` (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(jnp.ndarray)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `jnp.ndarray` (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.
    NÚlast_hidden_stateÚextract_featuresÚhidden_statesÚ
attentions)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚjnpÚndarrayÚ__annotations__r   r   r   r   r    © ó    úq/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/wav2vec2/modeling_flax_wav2vec2.pyr   r   ,   sW   … ñð, &*Ð�s—{‘{Ó)Ø$(Ð�c—k‘kÓ(Ø26€M�8˜E #§+¡+Ñ.Ñ/Ó6Ø/3€J�˜˜sŸ{™{Ñ+Ñ,Ô3r)   r   c                   óÔ   — e Zd ZU dZdZej                  ed<   dZej                  ed<   dZ	ej                  ed<   dZ
eeej                        ed<   dZeeej                        ed<   y)Ú FlaxWav2Vec2ForPreTrainingOutputa%  
    Output type of [`FlaxWav2Vec2ForPreTrainingOutput`], with potential hidden states and attentions.

    Args:
        loss (*optional*, returned when model is in train mode, `jnp.ndarray` 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 (`jnp.ndarray` 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 (`jnp.ndarray` 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(jnp.ndarray)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `jnp.ndarray` (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(jnp.ndarray)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `jnp.ndarray` (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.
    NÚprojected_statesÚprojected_quantized_statesÚcodevector_perplexityr   r    )r!   r"   r#   r$   r-   r%   r&   r'   r.   r/   r   r   r   r    r(   r)   r*   r,   r,   J   sf   … ñð4 %)Ð�c—k‘kÓ(Ø.2Ð §¡Ó2Ø)-Ð˜3Ÿ;™;Ó-Ø26€M�8˜E #§+¡+Ñ.Ñ/Ó6Ø/3€J�˜˜sŸ{™{Ñ+Ñ,Ô3r)   r,   ÚshapeÚ	mask_probÚmask_lengthÚattention_maskÚ	min_masksÚreturnc                 óŠ  — | \  }}|dk  rt        d«      ‚||kD  rt        d|› d|› d�«      ‚t        ||z  |z  t        j                  j	                  d«      j                  «       z   «      }t        ||«      }||z  |kD  r||z  }t        j                  ||ft        ¬«      }t        j                  t        |«      D �	cg c]=  }	t        j                  j                  t        j                  ||dz
  z
  «      |d¬«      ‘Œ? c}	«      }
t        j                  |
d	d	…d	d	…d	f   |||f«      }
|
j                  |||z  «      }
t        j                  |«      d	d	d	d	…f   }t        j                  ||||f«      j                  |||z  «      }|
|z   }
t        j                  ||
dd
«       |�t        j                   ||d«      }|S c c}	w )aw  
    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.
            should be of size 2 where first element is batch size and 2nd is timesteps
        mask_prob:
            probability for each token to be chosen as start of the span to be masked. this will be multiplied by
            number of timesteps divided by length of mask span to mask approximately this percentage of all elements.
            however due to overlaps, the actual number will be smaller (unless no_overlap is True)
        mask_length: size of the mask
        min_masks: minimum number of masked spans

    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`: ú`©ÚdtypeF)ÚreplaceNéÿÿÿÿ)Ú
ValueErrorÚintÚnpÚrandomÚrandÚitemÚmaxÚzerosÚboolÚarrayÚrangeÚchoiceÚarangeÚbroadcast_toÚreshapeÚput_along_axisÚwhere)r0   r1   r2   r3   r4   Ú
batch_sizeÚsequence_lengthÚnum_masked_spansÚspec_aug_maskÚ_Úspec_aug_mask_idxsÚoffsetss               r*   Ú_compute_mask_indicesrT   m   sì  € ð. #(Ñ€J�à�Q‚ÜÐAÓBÐBà�_Ò$ÜØ]Ð^iÐ]jð k#Ø#2Ð"3°1ð6ó
ð 	
ô ˜9 Ñ6¸ÑDÄrÇyÁyÇ~Á~ÐVWÓGX×G]ÑG]ÓG_Ñ_Ó`ÐÜÐ+¨YÓ7Ðð ˜+Ñ%¨Ò7Ø*¨kÑ9Ðô —H‘H˜j¨/Ð:Ä$ÔG€Mô Ÿ™ô ˜:Ó&ö	
àô �I‰I×ÑœRŸY™Y ¸+È¹/Ñ'JÓKÐM]ÐglÐÕmò	
óÐô Ÿ™Ð);ºAºqÀ$¸JÑ)GÈ*ÐVfÐhsÐItÓuÐØ+×3Ñ3°JÐ@PÐS^Ñ@^Ó_Ðä�i‰i˜Ó$ T¨4² ]Ñ3€GÜ�o‰o˜g¨
Ð4DÀkÐ'RÓS×[Ñ[ØÐ$ {Ñ2ó€Gð ,¨gÑ5Ðô ×Ñ�mÐ%7¸¸BÔ?àÐ!äŸ™ °ÀÓFˆàÐùò/	
s   Â>AG Úfeatures_shapeÚnum_negativesc                 ó8  — | \  }}}|dk  rt        d|||f› d�«      ‚g }t        |«      D ]V  }|�||   j                  «       dz
  n|dz
  }t        j                  j                  d|||z  f¬«      }	|j                  |	«       ŒX t        j                  |t        j                  ¬«      }t        j                  t        j                  |«      dd…df   ||f«      j                  «       }
|||
k\  xx   dz  cc<   t        d|«      D ]  }||xx   ||z  z  cc<   Œ |S )z>
    Sample `num_negatives` vectors from feature vectors.
    r   zl`features should have `sequence_length` > 1, but are of shape (batch_size, sequence_length, hidden_size) = (ú).Nr   )Úsizer8   )r<   rF   Úsumr>   r?   ÚrandintÚappendÚasarrayÚint32rI   rH   Úflatten)rU   rV   r3   rM   rN   Úhidden_sizeÚsampled_negative_indicesÚ	batch_idxÚhighÚsampled_indices_sliceÚfeature_indicess              r*   Ú_sample_negative_indicesrf   ¶   sQ  € ð 0>Ñ,€J� Ø˜!ÒÜð=Ø=GÈÐZeÐ=eÐ<fÐfhðjó
ð 	
ð  "ÐÜ˜:Ó&ò ?ˆ	Ø6DÐ6Pˆ~˜iÑ(×,Ñ,Ó.°Ò2ÐVeÐhiÑViˆÜ "§	¡	× 1Ñ 1°!°TÀÐQ`ÑA`Ð@bÐ 1Ó cÐØ ×'Ñ'Ð(=Õ>ð?ô
  "Ÿz™zÐ*BÌ"Ï(É(ÔSÐô —o‘o¤b§i¡i°Ó&@ÂÀDÀÑ&IÈOÐ]jÐKkÓl×tÑtÓv€Oð Ð5¸ÑHÓIÈQÑNÓIô ˜1˜jÓ)ò Kˆ	Ø  Ó+¨y¸?Ñ/JÑJÔ+ðKð $Ð#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 [`FlaxPreTrainedModel`]. Check the superclass documentation for the generic methods the
    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
    etc.)

    This model is also a Flax Linen
    [flax.nn.Module](https://flax.readthedocs.io/en/latest/_autosummary/flax.nn.module.html) subclass. Use it as a
    regular Flax Module and refer to the Flax documentation for all matter related to general usage and behavior.

    Finally, this model supports inherent JAX features such as:

    - [Just-In-Time (JIT) compilation](https://jax.readthedocs.io/en/latest/jax.html#just-in-time-compilation-jit)
    - [Automatic Differentiation](https://jax.readthedocs.io/en/latest/jax.html#automatic-differentiation)
    - [Vectorization](https://jax.readthedocs.io/en/latest/jax.html#vectorization-vmap)
    - [Parallelization](https://jax.readthedocs.io/en/latest/jax.html#parallelization-pmap)

    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 [`~FlaxPreTrainedModel.from_pretrained`] method to load the model weights.
        dtype (`jax.numpy.dtype`, *optional*, defaults to `jax.numpy.float32`):
            The data type of the computation. Can be one of `jax.numpy.float32`, `jax.numpy.float16` (on GPUs) and
            `jax.numpy.bfloat16` (on TPUs).

            This can be used to enable mixed-precision training or half-precision inference on GPUs or TPUs. If
            specified all the computation will be performed with the given `dtype`.

            **Note that this only specifies the dtype of the computation and does not influence the dtype of model
            parameters.**

            If you wish to change the dtype of the model parameters, see [`~FlaxPreTrainedModel.to_fp16`] and
            [`~FlaxPreTrainedModel.to_bf16`].
a‡	  
    Args:
        input_values (`jnp.ndarray` 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 `jnp.ndarray`. See [`Wav2Vec2Processor.__call__`] for details.
        attention_mask (`jnp.ndarray` 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) .. warning:: `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.
        mask_time_indices (`jnp.ndarray` 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.
        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.
c                   óh   — e Zd ZU eed<   dZeed<   ej                  Z	ej                  ed<   d„ Z
d„ Zy)ÚFlaxWav2Vec2LayerNormConvLayerÚconfigr   Úlayer_idr9   c           	      ó  — | j                   dkD  r#| j                  j                  | j                      nd| _        | j                  j                  | j                      | _        t        j                  | j                  j                  | j                      | j                  j                  | j                      f| j                  j                  | j                      f| j                  j                  t        j
                  j                  j                  «       d| j                  ¬«      | _        t        j                  | j                  j                   | j                  ¬«      | _        t$        | j                  j&                     | _        y )Nr   r   ÚVALID)ÚfeaturesÚkernel_sizeÚstridesÚuse_biasÚkernel_initÚpaddingr9   ©Úepsilonr9   )rj   ri   Úconv_dimÚin_conv_dimÚout_conv_dimÚnnÚConvÚconv_kernelÚconv_strideÚ	conv_biasÚjaxÚinitializersÚ	he_normalr9   ÚconvÚ	LayerNormÚlayer_norm_epsÚ
layer_normr   Úfeat_extract_activationÚ
activation©Úselfs    r*   Úsetupz$FlaxWav2Vec2LayerNormConvLayer.setup&  s  € ØBFÇ-Á-ÐRSÒBS˜4Ÿ;™;×/Ñ/°·±Ò>ÐYZˆÔØ ŸK™K×0Ñ0°·±Ñ?ˆÔä—G‘GØ—[‘[×)Ñ)¨$¯-©-Ñ8ØŸ™×0Ñ0°·±Ñ?ÐAØ—[‘[×,Ñ,¨T¯]©]Ñ;Ð=Ø—[‘[×*Ñ*ÜŸ™×+Ñ+×5Ñ5Ó7ØØ—*‘*ô
ˆŒ	ô Ÿ,™,¨t¯{©{×/IÑ/IÐQU×Q[ÑQ[Ô\ˆŒÜ  §¡×!DÑ!DÑEˆ�r)   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S ©N)r€   rƒ   r…   ©r‡   r   s     r*   Ú__call__z'FlaxWav2Vec2LayerNormConvLayer.__call__6  s2   € ØŸ	™	 -Ó0ˆØŸ™¨Ó6ˆØŸ™¨Ó6ˆØÐr)   N)r!   r"   r#   r   r'   rj   r=   r%   Úfloat32r9   rˆ   rŒ   r(   r)   r*   rh   rh   !  s/   … ØÓØ€HˆcÓØ—{‘{€Eˆ3�9‰9Ó"òFó r)   rh   c                   ó`   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	d„ Z
y)ÚFlaxConvWithWeightNormri   r9   c                 óf  ‡ — t        j                  ‰ j                  j                  ‰ j                  j                  ft
        j                   j                  j                  «       d‰ j                  j                  ‰ j                  ¬«      ‰ _
        ‰ j                  j                  ‰ j                  j                  ‰ j                  j                  z  ‰ j                  j                  d   f}‰ j                  dt
        j                   j                  j                  «       |«      ‰ _        ‰ j                  dˆ fd„«      ‰ _        ‰ j                  dt
        j                   j                  j"                  ‰ j                  j                  f«      ‰ _        ‰ j                  j                  d   dz  ‰ _        y )	Nrl   )rm   rn   rq   rr   Úfeature_group_countr9   r   Úweight_vÚweight_gc                 ój   •— t         j                  j                  ‰j                  d¬«      d d d d …f   S ©N)r   r   ©Úaxis)r%   ÚlinalgÚnormr’   )rQ   r‡   s    €r*   ú<lambda>z.FlaxConvWithWeightNorm.setup.<locals>.<lambda>P  s,   ø€ ¼¿¹¿¹ÈÏÉÐ]c¸Ó9dÐeiÐkoÒqrÐerÑ9s€ r)   Úbiasé   )rx   ry   ri   r`   Únum_conv_pos_embeddingsr}   r~   r   Únum_conv_pos_embedding_groupsr9   r€   rm   r‘   rn   Úparamr’   r“   rC   r›   Úprev_padding)r‡   Úweight_shapes   ` r*   rˆ   zFlaxConvWithWeightNorm.setupA  s+  ø€ Ü—G‘GØ—[‘[×,Ñ,ØŸ™×<Ñ<Ð>ÜŸ™×+Ñ+×5Ñ5Ó7ØØ $§¡× IÑ IØ—*‘*ô
ˆŒ	ð �I‰I×ÑØ�I‰I×Ñ $§)¡)×"?Ñ"?Ñ?Ø�I‰I×!Ñ! !Ñ$ð
ˆð
 Ÿ
™
 :¬s¯v©v×/BÑ/B×/LÑ/LÓ/NÐP\Ó]ˆŒØŸ
™
 :Ó/sÓtˆŒØ—J‘J˜v¤s§v¡v×':Ñ':×'@Ñ'@À4Ç9Á9×CUÑCUÐBWÓXˆŒ	Ø ŸI™I×1Ñ1°!Ñ4¸Ñ9ˆÕr)   c                 óì   — t         j                  j                  | j                  d¬«      d d d d …f   }t        j                  | j                  |«      }t        j
                  || j                  «      }|S r•   )r%   r˜   r™   r’   ÚdivideÚmultiplyr“   )r‡   Úweight_v_normÚnormed_weight_vÚnormed_kernels       r*   Ú_get_normed_weightsz*FlaxConvWithWeightNorm._get_normed_weightsT  sV   € ÜŸ
™
Ÿ™¨¯©¸F˜ÓCÀDÈ$ÒPQÀMÑRˆÜŸ*™* T§]¡]°MÓBˆÜŸ™ _°d·m±mÓDˆØÐr)   c                 óî   — | j                  «       }t        j                  |d| j                  | j                  fdf«      }| j                  j                  d|j                  | j                  dœi|«      }|S )N)r   r   Úparams)Úkernelr›   )r¨   r%   Úpadr    r€   ÚapplyÚTr›   )r‡   r   r«   s      r*   rŒ   zFlaxConvWithWeightNorm.__call__Z  si   € Ø×)Ñ)Ó+ˆÜŸ™ °¸×9JÑ9JÈD×L]ÑL]Ð8^Ð`fÐ/gÓhˆØŸ	™	Ÿ™¨¸f¿h¹hÐPT×PYÑPYÑ3ZÐ([Ð]jÓkˆØÐr)   N)r!   r"   r#   r   r'   r%   r�   r9   rˆ   r¨   rŒ   r(   r)   r*   r�   r�   =  s)   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò:ò&ór)   r�   c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)Ú#FlaxWav2Vec2PositionalConvEmbeddingri   r9   c                 óê   — t        | j                  | j                  ¬«      | _        t        | j                  j
                     | _        | j                  j                  dz  dk(  rd| _        y d| _        y )Nr8   rœ   r   r   )	r�   ri   r9   r€   r   r„   r…   r�   Únum_pad_remover†   s    r*   rˆ   z)FlaxWav2Vec2PositionalConvEmbedding.setupe  sT   € Ü*¨4¯;©;¸d¿j¹jÔIˆŒ	Ü  §¡×!DÑ!DÑEˆŒØ#'§;¡;×#FÑ#FÈÑ#JÈaÒ#O˜aˆÕÐUVˆÕr)   c                 óÞ   — |j                  d«      }| j                  |«      }| j                  dkD  r|d d …d | j                   …d d …f   }| j                  |«      }|j                  d«      }|S )N)r   r   rœ   r   )Ú	transposer€   r²   r…   r‹   s     r*   rŒ   z,FlaxWav2Vec2PositionalConvEmbedding.__call__j  sr   € Ø%×/Ñ/°	Ó:ˆàŸ	™	 -Ó0ˆà×Ñ Ò"Ø)ª!Ð-C°×0CÑ0CÐ/CÐ-CÂQÐ*FÑGˆMØŸ™¨Ó6ˆà%×/Ñ/°	Ó:ˆØÐr)   N©
r!   r"   r#   r   r'   r%   r�   r9   rˆ   rŒ   r(   r)   r*   r°   r°   a  s%   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òWó

r)   r°   c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)ÚFlaxConvLayersCollectionri   r9   c           
      ó†  — | j                   j                  dk(  r]t        | j                   j                  «      D �cg c].  }t	        | j                   |t        |«      | j                  ¬«      ‘Œ0 c}| _        y | j                   j                  dk(  rt        d«      ‚t        d| j                   j                  › d�«      ‚c c}w )NÚlayer)rj   Únamer9   ÚgroupzFAt the moment only ``config.feat_extact_norm == 'layer'`` is supportedz`config.feat_extract_norm` is z), but has to be one of ['group', 'layer'])
ri   Úfeat_extract_normrF   Únum_feat_extract_layersrh   Ústrr9   ÚlayersÚNotImplementedErrorr<   ©r‡   Úis     r*   rˆ   zFlaxConvLayersCollection.setup{  s©   € Ø�;‰;×(Ñ(¨GÒ3ô ˜tŸ{™{×BÑBÓCöàô /¨t¯{©{ÀQÌSÐQRËVÐ[_×[eÑ[eÖfòˆD�Kð �[‰[×*Ñ*¨gÒ5Ü%Ð&nÓoÐoäØ0°·±×1NÑ1NÐ0Oð Pð óð ùòs   »3B>c                 óP   — t        | j                  «      D ]  \  }} ||«      }Œ |S rŠ   )Ú	enumerater¿   )r‡   r   rÂ   Ú
conv_layers       r*   rŒ   z!FlaxConvLayersCollection.__call__‰  s.   € Ü& t§{¡{Ó3ò 	6‰MˆAˆzÙ& }Ó5‰Mð	6àÐr)   Nrµ   r(   r)   r*   r·   r·   w  s$   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òór)   r·   c                   ó`   — e Zd ZU dZeed<   ej                  Zej                  ed<   d„ Z	dd„Z
y)ÚFlaxWav2Vec2FeatureEncoderz.Construct the features from raw audio waveformri   r9   c                 óP   — t        | j                  | j                  ¬«      | _        y )Nr8   )r·   ri   r9   Úconv_layersr†   s    r*   rˆ   z FlaxWav2Vec2FeatureEncoder.setup•  s   € Ü3°D·K±KÀtÇzÁzÔRˆÕr)   c                 ó‚   — |d d …d d …d f   }| j                  |«      }|rt        j                  j                  |«      }|S rŠ   )rÉ   r}   r   Ústop_gradient)r‡   Úinput_valuesÚfreeze_feature_encoderr   s       r*   rŒ   z#FlaxWav2Vec2FeatureEncoder.__call__˜  s?   € Ø$¢Qª¨4 ZÑ0ˆØ×(Ñ(¨Ó7ˆÙ!ÜŸG™G×1Ñ1°-Ó@ˆMØÐr)   N)F)r!   r"   r#   r$   r   r'   r%   r�   r9   rˆ   rŒ   r(   r)   r*   rÇ   rÇ   �  s(   … Ù8àÓØ—{‘{€Eˆ3�9‰9Ó"òSôr)   rÇ   c                   ó\   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdd„Z	y)ÚFlaxWav2Vec2FeatureProjectionri   r9   c                 óÂ  — t        j                  | j                  j                  | j                  ¬«      | _        t        j                  | j                  j                  t        j                   j                  j                  | j                  j                  «      | j                  ¬«      | _        t        j                  | j                  j                  ¬«      | _        y )Nrs   ©rq   r9   ©Úrate)rx   r�   ri   r‚   r9   rƒ   ÚDenser`   r}   r~   ÚnormalÚinitializer_rangeÚ
projectionÚDropoutÚfeat_proj_dropoutÚdropoutr†   s    r*   rˆ   z#FlaxWav2Vec2FeatureProjection.setup¤  s‡   € ÜŸ,™,¨t¯{©{×/IÑ/IÐQU×Q[ÑQ[Ô\ˆŒÜŸ(™(Ø�K‰K×#Ñ#ÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQØ—*‘*ô
ˆŒô
 —z‘z t§{¡{×'DÑ'DÔEˆ�r)   c                 ót   — | j                  |«      }| j                  |«      }| j                  ||¬«      }||fS ©N©Údeterministic)rƒ   r×   rÚ   )r‡   r   rÞ   Únorm_hidden_statess       r*   rŒ   z&FlaxWav2Vec2FeatureProjection.__call__­  s>   € Ø!Ÿ_™_¨]Ó;ÐØŸ™Ð(:Ó;ˆØŸ™ ]À-˜ÓPˆØÐ0Ð0Ð0r)   N©Trµ   r(   r)   r*   rÏ   rÏ      s%   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òFô1r)   rÏ   c                   ó  — e Zd ZU eed<   eed<   eed<   dZeed<   dZe	ed<   e
j                  Ze
j                  ed<   dd„Zd„ Zd„ Z	 	 	 dde
j                   dee
j                      dee
j                      de	d	ee
j                      f
d„Zy
)ÚFlaxWav2Vec2Attentionri   Ú	embed_dimÚ	num_headsç        rÚ   Tr›   r9   r5   Nc           	      ór  — | j                   | j                  z  | _        | j                  | j                  z  | j                   k7  r&t        d| j                   › d| j                  › d�«      ‚t	        t
        j                  | j                   | j                  | j                  t        j
                  j                  j                  | j                  j                  «      ¬«      } |«        |«        |«       c| _        | _        | _         |«       | _        t        j$                  | j&                  ¬«      | _        y )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: rX   )rp   r9   rq   rÒ   )rã   rä   Úhead_dimr<   r   rx   rÔ   r›   r9   r}   r~   rÕ   ri   rÖ   Úq_projÚk_projÚv_projÚout_projrØ   rÚ   Údropout_layer)r‡   Údenses     r*   rˆ   zFlaxWav2Vec2Attention.setup¼  så   € ØŸ™¨$¯.©.Ñ8ˆŒØ�=‰=˜4Ÿ>™>Ñ)¨T¯^©^Ò;ÜØMÈdÏnÉnÐM]ð ^Ø—N‘NÐ# 2ð'óð ô
 Ü�H‰HØ�N‰NØ—Y‘YØ—*‘*ÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQô
ˆñ 16³¹»Á%Ã'Ð-ˆŒ�T”[ $¤+Ù›ˆŒäŸZ™Z¨T¯\©\Ô:ˆÕr)   c                 óp   — |j                  |j                  d d | j                  | j                  fz   «      S ©Nrœ   )rJ   r0   rä   rç   r‹   s     r*   Ú_split_headsz"FlaxWav2Vec2Attention._split_headsÑ  s5   € Ø×$Ñ$ ]×%8Ñ%8¸¸!Ð%<ÀÇÁÐPT×P]ÑP]Ð?^Ñ%^Ó_Ð_r)   c                 óZ   — |j                  |j                  d d | j                  fz   «      S rï   )rJ   r0   rã   r‹   s     r*   Ú_merge_headsz"FlaxWav2Vec2Attention._merge_headsÔ  s,   € Ø×$Ñ$ ]×%8Ñ%8¸¸!Ð%<ÀÇÁÐ?PÑ%PÓQÐQr)   r   Úkey_value_statesr3   rÞ   c                 óz  — | j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }|�t	        j
                  |d¬«      }|�°t        j                  |dkD  t	        j                  |j                  d«      j                  | j                  «      t	        j                  |j                  t	        j                  | j                  «      j                  «      j                  | j                  «      «      }nd}d}	|s | j                  dkD  r| j                  d«      }	t!        ||||	| j                  d|| j                  d¬«	      }
t	        j"                  d	|
|«      }| j%                  |«      }| j'                  |«      }||
fS )
z#Input shape: Batch x Time x ChannelN)éýÿÿÿéþÿÿÿr–   r   rå   rÚ   T)r›   Údropout_rngÚdropout_rateÚbroadcast_dropoutrÞ   r9   Ú	precisionz...hqk,...khd->...qhd)rè   ré   rê   rð   r%   Úexpand_dimsr   ÚselectÚfullr0   Úastyper9   ÚfinfoÚminrÚ   Úmake_rngr
   Úeinsumrò   rë   )r‡   r   ró   r3   rÞ   Úquery_statesÚ
key_statesÚvalue_statesÚattention_biasr÷   Úattn_weightsÚattn_outputs               r*   rŒ   zFlaxWav2Vec2Attention.__call__×  s�  € ð —{‘{ =Ó1ˆà—[‘[ Ó/ˆ
Ø—{‘{ =Ó1ˆà×(Ñ(¨Ó6ˆØ×&Ñ& zÓ2ˆ
Ø×(Ñ(¨Ó6ˆàÐ%Ü Ÿ_™_¨^À(ÔKˆNð Ð%ä ŸZ™ZØ Ñ"Ü—‘˜×-Ñ-¨sÓ3×:Ñ:¸4¿:¹:ÓFÜ—‘˜×-Ñ-¬s¯y©y¸¿¹Ó/D×/HÑ/HÓI×PÑPÐQU×Q[ÑQ[Ó\ó‰Nð "ˆNàˆÙ §¡°Ò!3ØŸ-™-¨	Ó2ˆKä4ØØØØ#ØŸ™Ø"Ø'Ø—*‘*Øô

ˆô —j‘jÐ!8¸,ÈÓUˆØ×'Ñ'¨Ó4ˆØ—m‘m KÓ0ˆà˜LÐ(Ð(r)   )r5   N)NNT)r!   r"   r#   r   r'   r=   rÚ   Úfloatr›   rD   r%   r�   r9   rˆ   rð   rò   r&   r   r   rŒ   r(   r)   r*   râ   râ   ´  s¨   … ØÓØƒNØƒNØ€GˆUÓØ€Dˆ$ÓØ—{‘{€Eˆ3�9‰9Ó"ó;ò*`òRð 37Ø04Ø"ñ5)à—{‘{ð5)ð # 3§;¡;Ñ/ð5)ð ! §¡Ñ-ð	5)ð
 ð5)ð 
ˆs�{‰{Ñ	ô5)r)   râ   c                   ó\   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdd„Z	y)ÚFlaxWav2Vec2FeedForwardri   r9   c                 ó\  — t        j                  | j                  j                  ¬«      | _        t        j
                  | j                  j                  t        j                   j                  j                  | j                  j                  «      | j                  ¬«      | _        t        | j                  j                  t        «      r#t         | j                  j                     | _        n| j                  j                  | _        t        j
                  | j                  j$                  t        j                   j                  j                  | j                  j                  «      | j                  ¬«      | _        t        j                  | j                  j(                  ¬«      | _        y )NrÒ   rÑ   )rx   rØ   ri   Úactivation_dropoutÚintermediate_dropoutrÔ   Úintermediate_sizer}   r~   rÕ   rÖ   r9   Úintermediate_denseÚ
isinstanceÚ
hidden_actr¾   r   Úintermediate_act_fnr`   Úoutput_denseÚhidden_dropoutÚoutput_dropoutr†   s    r*   rˆ   zFlaxWav2Vec2FeedForward.setup  s  € Ü$&§J¡J°D·K±K×4RÑ4RÔ$SˆÔ!ä"$§(¡(Ø�K‰K×)Ñ)ÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQØ—*‘*ô#
ˆÔô
 �d—k‘k×,Ñ,¬cÔ2Ü'-¨d¯k©k×.DÑ.DÑ'EˆDÕ$à'+§{¡{×'=Ñ'=ˆDÔ$äŸH™HØ�K‰K×#Ñ#ÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQØ—*‘*ô
ˆÔô
 !Ÿj™j¨d¯k©k×.HÑ.HÔIˆÕr)   c                 ó¸   — | j                  |«      }| j                  |«      }| j                  ||¬«      }| j                  |«      }| j	                  ||¬«      }|S rÜ   )r  r  r  r  r  ©r‡   r   rÞ   s      r*   rŒ   z FlaxWav2Vec2FeedForward.__call__'  sb   € Ø×/Ñ/°Ó>ˆØ×0Ñ0°Ó?ˆØ×1Ñ1°-È}Ð1Ó]ˆà×)Ñ)¨-Ó8ˆØ×+Ñ+¨MÈÐ+ÓWˆØÐr)   Nrà   rµ   r(   r)   r*   r  r    s%   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òJô(r)   r  c                   ó\   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdd„Z	y)Ú'FlaxWav2Vec2EncoderLayerStableLayerNormri   r9   c                 ó`  — t        | j                  | j                  j                  | j                  j                  | j                  j                  | j
                  ¬«      | _        t        j                  | j                  j                  ¬«      | _
        t        j                  | j                  j                  | j
                  ¬«      | _        t        | j                  | j
                  ¬«      | _        t        j                  | j                  j                  | j
                  ¬«      | _        y )N)ri   rã   rä   rÚ   r9   rÒ   rs   r8   )râ   ri   r`   Únum_attention_headsÚattention_dropoutr9   Ú	attentionrx   rØ   r  rÚ   r�   r‚   rƒ   r  Úfeed_forwardÚfinal_layer_normr†   s    r*   rˆ   z-FlaxWav2Vec2EncoderLayerStableLayerNorm.setup5  s½   € Ü.Ø—;‘;Ø—k‘k×-Ñ-Ø—k‘k×5Ñ5Ø—K‘K×1Ñ1Ø—*‘*ô
ˆŒô —z‘z t§{¡{×'AÑ'AÔBˆŒÜŸ,™,¨t¯{©{×/IÑ/IÐQU×Q[ÑQ[Ô\ˆŒÜ3°D·K±KÀtÇzÁzÔRˆÔÜ "§¡°T·[±[×5OÑ5OÐW[×WaÑWaÔ bˆÕr)   Nc                 óê   — |}| j                  |«      }| j                  |||¬«      \  }}| j                  ||¬«      }||z   }|| j                  | j	                  |«      |¬«      z   }|f}|r||fz  }|S )N)r3   rÞ   rÝ   )rƒ   r  rÚ   r  r   )r‡   r   r3   rÞ   Úoutput_attentionsÚattn_residualr  Úoutputss           r*   rŒ   z0FlaxWav2Vec2EncoderLayerStableLayerNorm.__call__B  sœ   € Ø%ˆØŸ™¨Ó6ˆØ&*§n¡nØ¨.Èð '5ó '
Ñ#ˆ�|ð Ÿ™ ]À-˜ÓPˆØ%¨Ñ5ˆØ%¨×(9Ñ(9Ø×!Ñ! -Ó0Àð ):ó )
ñ 
ˆð !Ð"ˆáØ˜�Ñ&ˆGàˆr)   )NTFrµ   r(   r)   r*   r  r  1  s%   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òcôr)   r  c            	       óx   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 	 d
de	de	de	de	fd	„Z
y)Ú1FlaxWav2Vec2EncoderLayerStableLayerNormCollectionri   r9   c           	      óÄ   — t        | j                  j                  «      D �cg c]-  }t        | j                  t	        |«      | j
                  ¬«      ‘Œ/ c}| _        y c c}w ©N)rº   r9   )rF   ri   Únum_hidden_layersr  r¾   r9   r¿   rÁ   s     r*   rˆ   z7FlaxWav2Vec2EncoderLayerStableLayerNormCollection.setupZ  sJ   € ô ˜4Ÿ;™;×8Ñ8Ó9ö
àô 4°D·K±KÄcÈ!ÃfÐTX×T^ÑT^Ö_ò
ˆ�ùò 
ó   ¢2ANrÞ   r"  Úoutput_hidden_statesÚreturn_dictc                 óü   — |rdnd }|rdnd }t        | j                  «      D ]*  \  }	}
|r||fz  } |
||||¬«      }|d   }|sŒ"||d   fz  }Œ, |r||fz  }|||f}|st        d„ |D «       «      S t        |||¬«      S )Nr(   )rÞ   r"  r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrŠ   r(   ©Ú.0Úvs     r*   ú	<genexpr>zMFlaxWav2Vec2EncoderLayerStableLayerNormCollection.__call__.<locals>.<genexpr>  ó   è ø€ Ò=˜q¨q©}œÑ=ùó   ‚Š©r   r   r    )rÄ   r¿   Útupler   )r‡   r   r3   rÞ   r"  r+  r,  Úall_attentionsÚall_hidden_statesrÂ   r¹   Úlayer_outputsr$  s                r*   rŒ   z:FlaxWav2Vec2EncoderLayerStableLayerNormCollection.__call__`  sÃ   € ñ  1™°dˆÙ"6™B¸DÐä! $§+¡+Ó.ò 	6‰HˆAˆuÙ#Ø! mÐ%5Ñ5Ð!á!Ø˜~¸]Ð^oôˆMð *¨!Ñ,ˆMâ Ø =°Ñ#3Ð"5Ñ5‘ð	6ñ  Ø -Ð!1Ñ1Ðà Ð"3°^ÐDˆáÜÑ= GÔ=Ó=Ð=ä"Ø+Ð;LÐYgô
ð 	
r)   ©NTFFT)r!   r"   r#   r   r'   r%   r�   r9   rˆ   rD   rŒ   r(   r)   r*   r&  r&  V  s]   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð Ø"Ø"'Ø%*Ø ñ#
ð ð	#
ð
  ð#
ð #ð#
ð ô#
r)   r&  c                   óf   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 	 dd„Z	y)Ú"FlaxWav2Vec2StableLayerNormEncoderri   r9   c                 ón  — t        | j                  | j                  ¬«      | _        t	        j
                  | j                  j                  | j                  ¬«      | _        t	        j                  | j                  j                  ¬«      | _
        t        | j                  | j                  ¬«      | _        y )Nr8   rs   rÒ   )r°   ri   r9   Úpos_conv_embedrx   r�   r‚   rƒ   rØ   r  rÚ   r&  r¿   r†   s    r*   rˆ   z(FlaxWav2Vec2StableLayerNormEncoder.setupŠ  sr   € ÜAÀ$Ç+Á+ÐUY×U_ÑU_Ô`ˆÔÜŸ,™,¨t¯{©{×/IÑ/IÐQU×Q[ÑQ[Ô\ˆŒÜ—z‘z t§{¡{×'AÑ'AÔBˆŒÜGÈÏÉÐ[_×[eÑ[eÔfˆ�r)   Nc                 óÈ  — |�?t        j                  t        j                  |d d …d d …d f   |j                  «      |d«      }| j	                  |«      }||z   }| j                  ||¬«      }| j                  |||||¬«      }| j                  |d   «      }	d }|r|d   }|d d |	fz   }|s#|	|f|r|dd  n|dd  z   }t        d„ |D «       «      S t        |	||j                  ¬«      S )	Nr   rÝ   )r"  r+  r,  r   r;   rœ   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrŠ   r(   r/  s     r*   r2  z>FlaxWav2Vec2StableLayerNormEncoder.__call__.<locals>.<genexpr>¶  r3  r4  r5  )r%   rL   rI   r0   r>  rÚ   r¿   rƒ   r6  r   r    )
r‡   r   r3   rÞ   r"  r+  r,  Úposition_embeddingsr$  r   s
             r*   rŒ   z+FlaxWav2Vec2StableLayerNormEncoder.__call__�  s!  € ð Ð%äŸI™IÜ× Ñ  ²²1°d°
Ñ!;¸]×=PÑ=PÓQÐS`ÐbcóˆMð #×1Ñ1°-Ó@Ðà%Ð(;Ñ;ˆØŸ™ ]À-˜ÓPˆà—+‘+ØØØ/Ø!5Ø#ð ó 
ˆð !ŸO™O¨G°A©JÓ7Ðð ˆÙØ# A™JˆMØ)¨#¨2Ð.Ð2CÐ1EÑEˆMáØ(¨-Ð8ÑK_¸GÀAÀB¹KÐelÐmnÐmoÐepÑqˆGÜÑ= GÔ=Ó=Ð=ä"Ø/¸}ÐY`×YkÑYkô
ð 	
r)   r:  rµ   r(   r)   r*   r<  r<  †  s6   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ògð ØØØ"Øô*
r)   r<  c                   ór   — e Zd ZU dZeed<   ej                  Zej                  ed<   d„ Z	e
dd„«       Zd	d„Zy)
Ú!FlaxWav2Vec2GumbelVectorQuantizerz¬
    Vector quantization using gumbel softmax. See [CATEGORICAL REPARAMETERIZATION WITH
    GUMBEL-SOFTMAX](https://arxiv.org/pdf/1611.01144.pdf) for more information.
    ri   r9   c                 óØ  — | j                   j                  | _        | j                   j                  | _        | j                   j
                  | j                  z  dk7  r0t        d| j                   j
                  › d| j                  › d�«      ‚| j                  dt        j                  j                  j                  «       d| j                  | j                  z  | j                   j
                  | j                  z  f«      | _        t        j                  | j                  | j                  z  t        j                  j                  j                  d«      | j                  ¬«      | _        y )	Nr   z`config.codevector_dim z5 must be divisible by `config.num_codevector_groups` z for concatenationÚcodevectorsr   ç      ð?rÑ   )ri   Únum_codevector_groupsÚ
num_groupsÚnum_codevectors_per_groupÚnum_varsÚcodevector_dimr<   rŸ   r}   rx   r~   ÚuniformrE  rÔ   rÕ   r9   Úweight_projr†   s    r*   rˆ   z'FlaxWav2Vec2GumbelVectorQuantizer.setupÆ  s  € ØŸ+™+×;Ñ;ˆŒØŸ™×=Ñ=ˆŒà�;‰;×%Ñ%¨¯©Ñ7¸1Ò<ÜØ)¨$¯+©+×*DÑ*DÐ)Eð F3Ø37·?±?Ð2CÐCUðWóð ð  Ÿ:™:ØÜ�F‰F×Ñ×'Ñ'Ó)Ø�—‘ $§-¡-Ñ/°·±×1KÑ1KÈtÏÉÑ1^Ð_ó
ˆÔô
 Ÿ8™8Ø�O‰O˜dŸm™mÑ+ÜŸ™×+Ñ+×2Ñ2°3Ó7Ø—*‘*ô
ˆÕr)   Nc           	      óÚ  — |�„t        j                  |j                  «       d d …d d f   | 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>r;   )
r%   rI   r_   r0   rL   Ú
zeros_likerZ   ÚmeanÚexpÚlog)ÚprobsÚmaskÚmask_extendedÚmarginal_probsÚ
perplexitys        r*   Ú_compute_perplexityz5FlaxWav2Vec2GumbelVectorQuantizer._compute_perplexityÜ  sµ   € àÐÜ×,Ñ,¨T¯\©\«^ºA¸tÀT¸MÑ-JÈEÏKÉKÓXˆMÜ—I‘I˜m¨U´C·N±NÀ5Ó4IÓJˆEØ"ŸY™Y¨A˜YÓ.°·±³Ñ;‰Nà"ŸZ™Z¨Q˜ZÓ/ˆNä—W‘WœcŸg™g n´s·w±w¸~ÐPTÑ?TÓ7UÑ&UÐ\^Ô_Ð_Ó`×dÑdÓfˆ
ØÐr)   c                 óÄ  — |j                   \  }}}| j                  |«      }|j                  ||z  | j                  z  d«      }|sž| j	                  d«      }t
        j                  j                  ||j                   «      }	t        j                  ||	z   |z  «      }
t        j                  |j                  ||z  | j                  d«      d¬«      }| j                  ||«      }nt|j                  d¬«      }t
        j                  j                  ||j                   d   «      dz  }
|
j                  ||z  | j                  d«      }
| j                  |
|«      }|
j                  ||z  d«      }
t        j                  |
d¬«      | j                  z  }|j                  ||z  | j                  | j                   d«      }|j#                  d«      j                  ||d«      }||fS )Nr;   Úgumbelr–   rF  rö   )r0   rM  rJ   rH  r  r}   r?   rZ  rx   ÚsoftmaxrX  ÚargmaxÚone_hotr%   rû   rE  rJ  rZ   )r‡   r   Úmask_time_indicesrÞ   ÚtemperaturerM   rN   r`   Ú
gumbel_rngÚgumbelsÚcodevector_probsÚcodevector_soft_distrW  Úcodevector_idxÚcodevectors_per_grouprE  s                   r*   rŒ   z*FlaxWav2Vec2GumbelVectorQuantizer.__call__è  sÍ  € Ø3@×3FÑ3FÑ0ˆ
�O [ð ×(Ñ(¨Ó7ˆØ%×-Ñ-¨j¸?Ñ.JÈTÏ_É_Ñ.\Ð^`ÓaˆáàŸ™ xÓ0ˆJÜ—j‘j×'Ñ'¨
°M×4GÑ4GÓHˆGÜ!Ÿz™z¨=¸7Ñ+BÀkÑ*QÓRÐô $&§:¡:Ø×%Ñ% j°?Ñ&BÀDÇOÁOÐUWÓXÐ_aô$Ð ð ×1Ñ1Ð2FÐHYÓZ‰Jð +×1Ñ1°rÐ1Ó:ˆNÜ"Ÿv™vŸ~™~¨n¸m×>QÑ>QÐRTÑ>UÓVÐY\Ñ\ÐØ/×7Ñ7¸
À_Ñ8TÐVZ×VeÑVeÐgiÓjÐØ×1Ñ1Ð2BÐDUÓVˆJà+×3Ñ3°JÀÑ4PÐRTÓUÐä #§¡Ð0@ÀrÔ JÈT×M]ÑM]Ñ ]ÐØ+×3Ñ3°JÀÑ4PÐRV×RaÑRaÐcg×cpÑcpÐrtÓuˆØ!—o‘o bÓ)×1Ñ1°*¸oÈrÓRˆà˜JÐ&Ð&r)   rŠ   )NTr   )r!   r"   r#   r$   r   r'   r%   r�   r9   rˆ   ÚstaticmethodrX  rŒ   r(   r)   r*   rC  rC  ½  s?   … ñð
 ÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð, ò	ó ð	ô 'r)   rC  c                   ó\   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdd„Z	y)ÚFlaxWav2Vec2Adapterri   r9   c                 ó(  — | j                   j                  | j                   j                  k7  r±t        j                  | j                   j                  t
        j                  j                  j                  | j                   j                  «      | j                  ¬«      | _
        t        j                  | j                   j                  | j                  ¬«      | _        nd x| _
        | _        t        | j                   | j                  ¬«      | _        y )NrÑ   rs   r8   )ri   Úoutput_hidden_sizer`   rx   rÔ   r}   r~   rÕ   rÖ   r9   Úprojr�   r‚   Úproj_layer_normÚ#FlaxWav2Vec2AdapterLayersCollectionr¿   r†   s    r*   rˆ   zFlaxWav2Vec2Adapter.setup  s¯   € à�;‰;×)Ñ)¨T¯[©[×-DÑ-DÒDÜŸ™Ø—‘×.Ñ.ÜŸF™F×/Ñ/×6Ñ6°t·{±{×7TÑ7TÓUØ—j‘jôˆDŒIô
 $&§<¡<¸¿¹×8RÑ8RÐZ^×ZdÑZdÔ#eˆDÕ à/3Ð3ˆDŒI˜Ô,ä9¸$¿+¹+ÈTÏZÉZÔXˆ�r)   c                 óœ   — | j                   �.| j                  �"| j                  |«      }| j                  |«      }| j                  |«      }|S rŠ   )rk  rl  r¿   r  s      r*   rŒ   zFlaxWav2Vec2Adapter.__call__  sI   € à�9‰9Ð  T×%9Ñ%9Ð%EØ ŸI™I mÓ4ˆMØ ×0Ñ0°Ó?ˆMàŸ™ MÓ2ˆàÐr)   Nrà   rµ   r(   r)   r*   rh  rh    s%   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òYôr)   rh  c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)ÚFlaxWav2Vec2AdapterLayerri   r9   c           	      óP  — t        j                  d| j                  j                  z  | j                  j                  f| j                  j
                  fdt        j                   j                  j                  | j                  j                  «      | j                  ¬«      | _        y )Nrœ   ))r   r   )rm   rn   ro   rr   rq   r9   )rx   ry   ri   rj  Úadapter_kernel_sizeÚadapter_strider}   r~   rÕ   rÖ   r9   r€   r†   s    r*   rˆ   zFlaxWav2Vec2AdapterLayer.setup,  sp   € Ü—G‘GØ˜Ÿ™×7Ñ7Ñ7ØŸ™×8Ñ8Ð:Ø—[‘[×/Ñ/Ð1ØÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQØ—*‘*ô
ˆ�	r)   c                 óV   — | j                  |«      }t        j                  |d¬«      }|S )Nrœ   r–   )r€   rx   Úglur‹   s     r*   rŒ   z!FlaxWav2Vec2AdapterLayer.__call__6  s&   € ØŸ	™	 -Ó0ˆÜŸ™˜}°1Ô5ˆàÐr)   Nrµ   r(   r)   r*   rp  rp  (  s$   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ór)   rp  c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)rm  ri   r9   c           	      óÄ   — t        | j                  j                  «      D �cg c]-  }t        | j                  t	        |«      | j
                  ¬«      ‘Œ/ c}| _        y c c}w r(  )rF   ri   Únum_adapter_layersrp  r¾   r9   r¿   rÁ   s     r*   rˆ   z)FlaxWav2Vec2AdapterLayersCollection.setupA  sG   € ô ˜4Ÿ;™;×9Ñ9Ó:ö
àô % T§[¡[´s¸1³vÀTÇZÁZÖPò
ˆ�ùò 
r*  c                 ó8   — | j                   D ]
  } ||«      }Œ |S rŠ   )r¿   )r‡   r   rÅ   s      r*   rŒ   z,FlaxWav2Vec2AdapterLayersCollection.__call__G  s'   € ØŸ+™+ò 	6ˆJÙ& }Ó5‰Mð	6ð Ðr)   Nrµ   r(   r)   r*   rm  rm  =  s$   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ór)   rm  c                   ó¦  ‡ — e Zd ZU dZeZdZeed<   dZ	dZ
ej                  ed<   ddej                  d	fd
edededej"                  def
ˆ fd„Zddej*                  j,                  dededefd„Z ee«      	 	 	 	 	 	 	 	 	 ddedej*                  j,                  dedee   dee   dedee   fd„«       Z	 ddeej>                  ef   dee   fd„Z ˆ xZ!S ) ÚFlaxWav2Vec2PreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úwav2vec2Úbase_model_prefixrÌ   NÚmodule_class)r   i   r   Tri   Úinput_shapeÚseedr9   Ú_do_initc                 óZ   •—  | j                   d||dœ|¤Ž}t        ‰| �	  ||||||¬«       y )N)ri   r9   )r  r€  r9   r�  r(   )r~  ÚsuperÚ__init__)	r‡   ri   r  r€  r9   r�  ÚkwargsÚmoduleÚ	__class__s	           €r*   r„  z$FlaxWav2Vec2PreTrainedModel.__init__Y  s=   ø€ ð #�×"Ñ"ÐH¨&¸ÑHÀÑHˆÜ‰Ñ˜ °[ÀtÐSXÐckÐÕlr)   Úrngrª   r5   c                 ó¾  — t        j                  |d¬«      }t        j                  |«      }t        j                  j                  |d«      \  }}||dœ}| j                  j                  |||d¬«      d   }	|�dt        t        |	«      «      }	t        t        |«      «      }| j                  D ]
  }
|	|
   ||
<   Œ t        «       | _
        t        t        |«      «      S |	S )NÚi4r8   rœ   )rª   rÚ   F)r,  rª   )r%   rC   Ú	ones_liker}   r?   Úsplitr†  Úinitr   r	   Ú_missing_keysÚsetr   r   )r‡   rˆ  r  rª   rÌ   r3   Ú
params_rngr÷   ÚrngsÚrandom_paramsÚmissing_keys              r*   Úinit_weightsz(FlaxWav2Vec2PreTrainedModel.init_weightse  sÓ   € ä—y‘y °DÔ9ˆÜŸ™ |Ó4ˆÜ"%§*¡*×"2Ñ"2°3¸Ó":Ñˆ
�KØ$°Ñ=ˆàŸ™×(Ñ(¨¨|¸^ÐY^Ð(Ó_Ð`hÑiˆàÐÜ(¬°-Ó)@ÓAˆMÜ!¤(¨6Ó"2Ó3ˆFØ#×1Ñ1ò A�Ø&3°KÑ&@��{Ò#ðAä!$£ˆDÔÜœ.¨Ó0Ó1Ð1à Ð r)   r÷   Útrainr"  r+  rÍ   r,  c                 óÄ  — |�|n| j                   j                  }|�|n| j                   j                  }|
�|
n| j                   j                  }
|j                  \  }}|€t        j                  ||f«      }i }|�||d<   d|xs | j                  i}| j                  j                  |t        j                  |d¬«      t        j                  |d¬«      || |||	|
|¬«
      S )NrÚ   rª   Úf4r8   rŠ  ©r‘  ©ri   r"  r+  r,  r0   r%   Úonesrª   r†  r­   rE   )r‡   rÌ   r3   r^  rª   r÷   r•  r"  r+  rÍ   r,  rM   rN   r‘  Úinputss                  r*   rŒ   z$FlaxWav2Vec2PreTrainedModel.__call__x  sö   € ð 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆà&2×&8Ñ&8Ñ#ˆ
�OàÐ!Ü ŸX™X z°?Ð&CÓDˆNð ˆØÐ"Ø)ˆD�‰Oà˜FÒ1 d§k¡kÐ2ˆà�{‰{× Ñ ØÜ�I‰I�l¨$Ô/Ü�I‰I�n¨DÔ1ØØˆIØØ Ø"ØØð !ó 
ð 	
r)   Úinput_lengthsÚadd_adapterc                 ó<   — | j                   j                  ||¬«      S )N©r�  )r†  Ú _get_feat_extract_output_lengths)r‡   rœ  r�  s      r*   r   z<FlaxWav2Vec2PreTrainedModel._get_feat_extract_output_lengths¥  s   € ð �{‰{×;Ñ;¸MÐWbÐ;ÓcÐcr)   rŠ   )	NNNNFNNFN)"r!   r"   r#   r$   r   Úconfig_classr}  r¾   r'   Úmain_input_namer~  rx   ÚModuler%   r�   r   r=   r9   rD   r„  r}   r?   ÚPRNGKeyr   r”  r   ÚWAV2VEC2_INPUTS_DOCSTRINGÚdictr   rŒ   r   r&   r   Ú__classcell__)r‡  s   @r*   r{  r{  N  sq  ø… ñð
 "€LØ'Ð�sÓ'Ø$€OØ"€L�"—)‘)Ó"ð
 'ØØŸ;™;Øñ
màð
mð ð
mð ð	
mð
 �y‰yð
mð õ
mñ! §
¡
× 2Ñ 2ð !Àð !ÐPZð !Ðfpó !ñ& +Ð+DÓEð ØØØ*.ØØ,0Ø/3Ø',Ø&*ñ*
ð
 ð*
ð —Z‘Z×'Ñ'ð*
ð ð*
ð $ D™>ð*
ð ' t™nð*
ð !%ð*
ð ˜d‘^ò*
ó Fð*
ðZ UYñdØ" 3§;¡;°Ð#3Ñ4ðdØCKÈDÁ>÷dr)   r{  c                   óÊ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 	 	 	 dd„Z		 dde
ej                  ef   dee   fd„Z	 dd	ed
ej                  fd„Zy)ÚFlaxWav2Vec2Moduleri   r9   c                 óZ  — t        | j                  | j                  ¬«      | _        t	        | j                  | j                  ¬«      | _        | j                  dt        j                  j                  j                  «       | j                  j                  f«      | _        | j                  j                  r't        | j                  | j                  ¬«      | _        nt!        d«      ‚| j                  j"                  r't%        | j                  | j                  ¬«      | _        y d | _        y )Nr8   Úmasked_spec_embedzD``config.do_stable_layer_norm is False`` is currently not supported.)rÇ   ri   r9   Úfeature_extractorrÏ   Úfeature_projectionrŸ   r}   rx   r~   rL  r`   r«  Údo_stable_layer_normr<  ÚencoderrÀ   r�  rh  Úadapterr†   s    r*   rˆ   zFlaxWav2Vec2Module.setup¯  sÉ   € Ü!;¸D¿K¹KÈtÏzÉzÔ!ZˆÔÜ"?ÀÇÁÐSW×S]ÑS]Ô"^ˆÔØ!%§¡Ø¤§¡×!4Ñ!4×!<Ñ!<Ó!>ÀÇÁ×AXÑAXÐ@Zó"
ˆÔð �;‰;×+Ò+Ü=¸d¿k¹kÐQU×Q[ÑQ[Ô\ˆD�Lä%Ð&lÓmÐmàMQÏ[É[×MdÒMdÔ*¨4¯;©;¸d¿j¹jÔIˆ�Ðjnˆ�r)   Nc	           
      óL  — | j                  ||¬«      }	|�!| j                  |	j                  d   |d¬«      }| j                  |	|¬«      \  }
}	|�ot	        j
                  t	        j                  |d d …d d …d f   |
j                  «      t	        j                  | j                  d d d d …f   |
j                  «      |
«      }
| j                  |
|||||¬«      }|d   }
| j                  �| j                  |
«      }
|s
|
|	f|dd  z   S t        |
|	|j                  |j                  ¬«      S )	N)rÍ   r   FrŸ  rÝ   )r3   rÞ   r"  r+  r,  r   )r   r   r   r    )r¬  Ú"_get_feature_vector_attention_maskr0   r­  r%   rL   rI   r«  r¯  r°  r   r   r    )r‡   rÌ   r3   r^  rÞ   r"  r+  rÍ   r,  r   r   Úencoder_outputss               r*   rŒ   zFlaxWav2Vec2Module.__call__½  sU  € ð  ×1Ñ1°,ÐWmÐ1ÓnÐð Ð%à!×DÑDØ ×&Ñ& qÑ)¨>Àuð Eó ˆNð +/×*AÑ*AÐBRÐboÐ*AÓ*pÑ'ˆÐ'ØÐ(ÜŸI™IÜ× Ñ Ð!2²1²a¸°:Ñ!>À×@SÑ@SÓTÜ× Ñ  ×!7Ñ!7¸¸dÂA¸Ñ!FÈ×H[ÑH[Ó\ØóˆMð Ÿ,™,ØØ)Ø'Ø/Ø!5Ø#ð 'ó 
ˆð (¨Ñ*ˆà�<‰<Ð#Ø ŸL™L¨Ó7ˆMáØ!Ð#3Ð4°ÀqÀrÐ7JÑJÐJä*Ø+Ø-Ø)×7Ñ7Ø&×1Ñ1ô	
ð 	
r)   rœ  r�  c                 óT  — |€| j                   j                  n|}d„ }t        | j                   j                  | j                   j                  «      D ]  \  }} ||||«      }Œ |rBt        | j                   j                  «      D ]   } ||d| j                   j                  «      }Œ" |S )úH
        Computes the output length of the convolutional layers
        c                 ó   — | |z
  |z  dz   S ©Nr   r(   ©Úinput_lengthrn   Ústrides      r*   Ú_conv_out_lengthzMFlaxWav2Vec2Module._get_feat_extract_output_lengths.<locals>._conv_out_lengthú  ó   € ð ! ;Ñ.°6Ñ9¸AÑ=Ð=r)   r   ©ri   r�  Úziprz   r{   rF   rx  rs  ©r‡   rœ  r�  r»  rn   rº  rQ   s          r*   r   z3FlaxWav2Vec2Module._get_feat_extract_output_lengthsñ  ó¥   € ð 2=Ð1D�d—k‘k×-Ò-È+ˆò	>ô
 $' t§{¡{×'>Ñ'>ÀÇÁ×@WÑ@WÓ#Xò 	QÑˆK˜Ù,¨]¸KÈÓP‰Mð	Qñ Ü˜4Ÿ;™;×9Ñ9Ó:ò _�Ù 0°ÀÀ4Ç;Á;×C]ÑC]Ó ^‘ð_ð Ðr)   Úfeature_vector_lengthr3   c                 óØ  — |j                  d¬«      d d …df   }| j                  ||¬«      }|j                  d   }t        j                  ||f|j
                  ¬«      }|j                  t        j                  |j                  d   «      |dz
  f   j                  d«      }t        j                  t        j                  |d«      j                  d«      d«      j                  d«      }|S )Nr;   r–   rŸ  r   r8   r   rD   )Úcumsumr   r0   r%   rC   r9   ÚatrH   r�  Úfliprþ   )r‡   rÁ  r3   r�  Únon_padded_lengthsÚoutput_lengthsrM   s          r*   r²  z5FlaxWav2Vec2Module._get_feature_vector_attention_mask  sÚ   € ð
 ,×2Ñ2¸Ð2Ó;ºA¸r¸EÑBÐà×>Ñ>Ð?QÐ_jÐ>Ókˆà#×)Ñ)¨!Ñ,ˆ
äŸ™ JÐ0EÐ#FÈn×NbÑNbÔcˆð (×*Ñ*¬3¯:©:°n×6JÑ6JÈ1Ñ6MÓ+NÐP^ÐabÑPbÐ+bÑc×gÑgÐhiÓjˆÜŸ™¤#§(¡(¨>¸2Ó">×"EÑ"EÀbÓ"IÈ2ÓN×UÑUÐV\Ó]ˆØÐr)   ©NNTNNFNrŠ   )r!   r"   r#   r   r'   r%   r�   r9   rˆ   rŒ   r   r&   r=   r   rD   r   r²  r(   r)   r*   r©  r©  «  s‹   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òoð" ØØØØ!Ø$Øó2
ðj UYñØ" 3§;¡;°Ð#3Ñ4ðØCKÈDÁ>óð0 TXñØ%(ðØ:=¿+¹+ôr)   r©  zbThe bare Wav2Vec2 Model transformer outputting raw hidden-states without any specific head on top.c                   ó   — e Zd ZeZy)ÚFlaxWav2Vec2ModelN)r!   r"   r#   r©  r~  r(   r)   r*   rÊ  rÊ    s	   „ ð
 &�Lr)   rÊ  aJ  
    Returns:

    Example:

    ```python
    >>> from transformers import AutoProcessor, FlaxWav2Vec2Model
    >>> from datasets import load_dataset
    >>> import soundfile as sf

    >>> processor = AutoProcessor.from_pretrained("facebook/wav2vec2-large-lv60")
    >>> model = FlaxWav2Vec2Model.from_pretrained("facebook/wav2vec2-large-lv60")


    >>> def map_to_array(batch):
    ...     speech, _ = sf.read(batch["file"])
    ...     batch["speech"] = speech
    ...     return batch


    >>> ds = load_dataset("hf-internal-testing/librispeech_asr_dummy", "clean", split="validation")
    >>> ds = ds.map(map_to_array)

    >>> input_values = processor(
    ...     ds["speech"][0], sampling_rate=16_000, return_tensors="np"
    ... ).input_values  # Batch size 1
    >>> hidden_states = model(input_values).last_hidden_state
    ```
)Úoutput_typer¡  c                   ó¢   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 	 	 	 d	d„Z		 d
de
ej                  ef   dee   fd„Zy)ÚFlaxWav2Vec2ForCTCModuleri   r9   c                 óš  — t        | j                  | j                  ¬«      | _        t	        j
                  | j                  j                  ¬«      | _        t	        j                  | j                  j                  t        j                  j                  j                  | j                  j                  «      | j                  ¬«      | _        y )Nr8   rÒ   rÑ   )r©  ri   r9   r|  rx   rØ   Úfinal_dropoutrÚ   rÔ   Ú
vocab_sizer}   r~   rÕ   rÖ   Úlm_headr†   s    r*   rˆ   zFlaxWav2Vec2ForCTCModule.setupN  sx   € Ü*¨4¯;©;¸d¿j¹jÔIˆŒÜ—z‘z t§{¡{×'@Ñ'@ÔAˆŒÜ—x‘xØ�K‰K×"Ñ"ÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQØ—*‘*ô
ˆ�r)   Nc	           
      óà   — | j                  ||||||||¬«      }	|	d   }
| j                  |
|¬«      }
| j                  |
«      }|s	|f|	dd  z   S t        ||	j                  |	j
                  ¬«      S )N)r3   r^  rÞ   r"  r+  rÍ   r,  r   rÝ   rœ   )Úlogitsr   r    )r|  rÚ   rÑ  r   r   r    )r‡   rÌ   r3   r^  rÞ   r"  r+  rÍ   r,  r$  r   rÓ  s               r*   rŒ   z!FlaxWav2Vec2ForCTCModule.__call__W  sŒ   € ð —-‘-ØØ)Ø/Ø'Ø/Ø!5Ø#9Ø#ð  ó 	
ˆð   ™
ˆØŸ™ ]À-˜ÓPˆà—‘˜mÓ,ˆáØ�9˜w q r˜{Ñ*Ð*ä!¨¸w×?TÑ?TÐah×asÑasÔtÐtr)   rœ  r�  c                 óT  — |€| j                   j                  n|}d„ }t        | j                   j                  | j                   j                  «      D ]  \  }} ||||«      }Œ |rBt        | j                   j                  «      D ]   } ||d| j                   j                  «      }Œ" |S )rµ  c                 ó   — | |z
  |z  dz   S r·  r(   r¸  s      r*   r»  zSFlaxWav2Vec2ForCTCModule._get_feat_extract_output_lengths.<locals>._conv_out_length‚  r¼  r)   r   r½  r¿  s          r*   r   z9FlaxWav2Vec2ForCTCModule._get_feat_extract_output_lengthsw  s¥   € ð 2=Ð1D�d—k‘k×-Ò-È+ˆò	>ô
 $' t§{¡{×'>Ñ'>ÀÇÁ×@WÑ@WÓ#Xò 	QÑˆK˜Ù,¨]¸KÈÓP‰Mð	Qñ Ü˜4Ÿ;™;×9Ñ9Ó:ò _�Ù 0°ÀÀ4Ç;Á;×C]ÑC]Ó ^‘ð_ð Ðr)   rÈ  rŠ   )r!   r"   r#   r   r'   r%   r�   r9   rˆ   rŒ   r   r&   r=   r   rD   r   r(   r)   r*   rÍ  rÍ  J  sk   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð ØØØØ!Ø$ØóuðF '+ñà˜SŸ[™[¨#Ð-Ñ.ðð ˜d‘^ôr)   rÍ  zfWav2Vec2 Model with a `language modeling` head on top for Connectionist Temporal Classification (CTC).c                   ó   — e Zd ZeZy)ÚFlaxWav2Vec2ForCTCN)r!   r"   r#   rÍ  r~  r(   r)   r*   r×  r×  ‘  s	   „ ð
 ,�Lr)   r×  a  
    Returns:

    Example:

    ```python
    >>> import jax.numpy as jnp
    >>> from transformers import AutoProcessor, FlaxWav2Vec2ForCTC
    >>> from datasets import load_dataset
    >>> import soundfile as sf

    >>> processor = AutoProcessor.from_pretrained("facebook/wav2vec2-large-960h-lv60")
    >>> model = FlaxWav2Vec2ForCTC.from_pretrained("facebook/wav2vec2-large-960h-lv60")


    >>> def map_to_array(batch):
    ...     speech, _ = sf.read(batch["file"])
    ...     batch["speech"] = speech
    ...     return batch


    >>> ds = load_dataset("hf-internal-testing/librispeech_asr_dummy", "clean", split="validation")
    >>> ds = ds.map(map_to_array)

    >>> input_values = processor(
    ...     ds["speech"][0], sampling_rate=16_000, return_tensors="np"
    ... ).input_values  # Batch size 1
    >>> logits = model(input_values).logits
    >>> predicted_ids = jnp.argmax(logits, axis=-1)

    >>> transcription = processor.decode(predicted_ids[0])
    >>> # should give:  "A MAN SAID TO THE UNIVERSE SIR I EXIST"
    ```
c                   ó®   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 	 	 	 	 dde	de
fd„Z	 ddeej                  e	f   d	ee
   fd
„Zy)Ú FlaxWav2Vec2ForPreTrainingModuleri   r9   c                 óÐ  — t        | j                  | j                  ¬«      | _        t	        j
                  | j                  j                  «      | _        t        | j                  | j                  ¬«      | _	        t	        j                  | j                  j                  t        j                  j                  j                  | j                  j                  «      | j                  ¬«      | _        t	        j                  | j                  j                  t        j                  j                  j                  | j                  j                  «      | j                  ¬«      | _        y )Nr8   rÑ   )r©  ri   r9   r|  rx   rØ   Úfeat_quantizer_dropoutÚdropout_featuresrC  Ú	quantizerrÔ   Úproj_codevector_dimr}   r~   rÕ   rÖ   Ú	project_qÚproject_hidr†   s    r*   rˆ   z&FlaxWav2Vec2ForPreTrainingModule.setupÇ  sØ   € Ü*¨4¯;©;¸d¿j¹jÔIˆŒÜ "§
¡
¨4¯;©;×+MÑ+MÓ NˆÔä:¸4¿;¹;ÈdÏjÉjÔYˆŒÜŸ™Ø�K‰K×+Ñ+ÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQØ—*‘*ô
ˆŒô
 Ÿ8™8Ø�K‰K×+Ñ+ÜŸ™×+Ñ+×2Ñ2°4·;±;×3PÑ3PÓQØ—*‘*ô
ˆÕr)   NÚgumbel_temperaturerÞ   c
           
      óp  — |	�|	n| j                   j                  }	| j                  ||||||||	¬«      }
| j                  |
d   «      }| j	                  |
d   |¬«      }| j                  ||||¬«      \  }}| j                  |«      }|	s|||f|
dd z   S t        ||||
j                  |
j                  ¬«      S )	zC
        Returns:

        Example:

        ```python

        ```N)r3   r"  r+  r^  rÞ   rÍ   r,  r   r   rÝ   )rÞ   r_  rœ   )r-   r.   r/   r   r    )
ri   Úuse_return_dictr|  rà  rÜ  rÝ  rß  r,   r   r    )r‡   rÌ   r3   r^  rá  rÞ   r"  r+  rÍ   r,  r$  Útransformer_featuresr   Úquantized_featuresr/   s                  r*   rŒ   z)FlaxWav2Vec2ForPreTrainingModule.__call__×  sù   € ð* &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—-‘-ØØ)Ø/Ø!5Ø/Ø'Ø#9Ø#ð  ó 	
ˆð  $×/Ñ/°¸±
Ó;Ðð  ×0Ñ0°¸±È=Ð0ÓYÐØ48·N±NØÐ/¸}ÐZlð 5Có 5
Ñ1ÐÐ1ð "Ÿ^™^Ð,>Ó?ÐáØ(Ð*<Ð>SÐTÐW^Ð_`Ð_aÐWbÑbÐbä/Ø1Ø'9Ø"7Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r)   rœ  r�  c                 óT  — |€| j                   j                  n|}d„ }t        | j                   j                  | j                   j                  «      D ]  \  }} ||||«      }Œ |rBt        | j                   j                  «      D ]   } ||d| j                   j                  «      }Œ" |S )rµ  c                 ó   — | |z
  |z  dz   S r·  r(   r¸  s      r*   r»  z[FlaxWav2Vec2ForPreTrainingModule._get_feat_extract_output_lengths.<locals>._conv_out_length  r¼  r)   r   r½  r¿  s          r*   r   zAFlaxWav2Vec2ForPreTrainingModule._get_feat_extract_output_lengths  rÀ  r)   )NNr   TNNFNrŠ   )r!   r"   r#   r   r'   r%   r�   r9   rˆ   r=   rD   rŒ   r   r&   r   r   r(   r)   r*   rÙ  rÙ  Ã  s�   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð& ØØ"#Ø"ØØ!Ø$Øñ5
ð
  ð5
ð ó5
ðp UYñØ" 3§;¡;°Ð#3Ñ4ðØCKÈDÁ>ôr)   rÙ  z5Wav2Vec2 Model with a quantizer and `VQ` head on top.c                   óÌ   — e Zd ZeZ ee«      	 	 	 	 	 	 	 	 	 	 	 ddedede	j                  j                  de	j                  j                  dedee   dee   d	ed
ee   fd„«       Zy)ÚFlaxWav2Vec2ForPreTrainingNrá  rª   r÷   r`  r•  r"  r+  rÍ   r,  c                 óÔ  — |	�|	n| j                   j                  }	|
�|
n| j                   j                  }
|�|n| j                   j                  }|j                  \  }}|€t        j                  ||f«      }i }|�||d<   |�||d<   d|xs | j                  i}| j                  j                  |t        j                  |d¬«      t        j                  |d¬«      ||| |	|
|||¬«      S )NrÚ   rZ  rª   r—  r8   rŠ  r˜  r™  )r‡   rÌ   r3   r^  rá  rª   r÷   r`  r•  r"  r+  rÍ   r,  rM   rN   r‘  r›  s                    r*   rŒ   z#FlaxWav2Vec2ForPreTraining.__call__*  s  € ð" 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆà&2×&8Ñ&8Ñ#ˆ
�OàÐ!Ü ŸX™X z°?Ð&CÓDˆNð ˆØÐ"Ø)ˆD�‰OàÐ!Ø'ˆD�‰Nà˜FÒ1 d§k¡kÐ2ˆà�{‰{× Ñ ØÜ�I‰I�l¨$Ô/Ü�I‰I�n¨DÔ1ØØØˆIØØ Ø"ØØð !ó 
ð 	
r)   )NNr   NNNFNNFN)r!   r"   r#   rÙ  r~  r   r¥  r=   r¦  r}   r?   r¤  rD   r   rŒ   r(   r)   r*   ré  ré  &  s½   „ à3€Lá*Ð+DÓEð
 ØØ"#ØØ*.Ø)-ØØ,0Ø/3Ø',Ø&*ñ0
ð
  ð0
ð ð0
ð —Z‘Z×'Ñ'ð0
ð —J‘J×&Ñ&ð0
ð ð0
ð $ D™>ð0
ð ' t™nð0
ð !%ð0
ð ˜d‘^ò0
ó Fñ0
r)   ré  a•  
    Returns:

    Example:

    ```python
    >>> import optax
    >>> import numpy as np
    >>> import jax.numpy as jnp
    >>> from transformers import AutoFeatureExtractor, FlaxWav2Vec2ForPreTraining
    >>> from transformers.models.wav2vec2.modeling_flax_wav2vec2 import _compute_mask_indices
    >>> from datasets import load_dataset
    >>> import soundfile as sf

    >>> feature_extractor = AutoFeatureExtractor.from_pretrained("facebook/wav2vec2-large-lv60")
    >>> model = FlaxWav2Vec2ForPreTraining.from_pretrained("facebook/wav2vec2-large-lv60")


    >>> def map_to_array(batch):
    ...     speech, _ = sf.read(batch["file"])
    ...     batch["speech"] = speech
    ...     return batch


    >>> ds = load_dataset("hf-internal-testing/librispeech_asr_dummy", "clean", split="validation")
    >>> ds = ds.map(map_to_array)

    >>> input_values = feature_extractor(ds["speech"][0], return_tensors="np").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)
    >>> mask_time_indices = _compute_mask_indices((batch_size, sequence_length), mask_prob=0.2, mask_length=2)

    >>> outputs = model(input_values, mask_time_indices=mask_time_indices)

    >>> # compute cosine similarity between predicted (=projected_states) and target (=projected_quantized_states)
    >>> cosine_sim = optax.cosine_similarity(outputs.projected_states, outputs.projected_quantized_states)

    >>> # show that cosine similarity is much higher than random
    >>> assert np.asarray(cosine_sim)[mask_time_indices].mean() > 0.5
    ```
)r×  ré  rÊ  r{  )Nr   rŠ   )Rr$   Ú	functoolsr   Útypingr   r   r   ÚflaxÚ
flax.linenÚlinenrx   r}   Ú	jax.numpyÚnumpyr%   r>   Úflax.core.frozen_dictr   r   r	   Úflax.linen.attentionr
   Úflax.traverse_utilr   r   r   Úmodeling_flax_outputsr   r   Úmodeling_flax_utilsr   r   r   r   Úutilsr   r   r   r   Úconfiguration_wav2vec2r   Ú
get_loggerr!   ÚloggerÚstructÚ	dataclassr   r,   r=   r	  r&   rT   rf   ÚWAV2VEC2_START_DOCSTRINGr¥  r£  rh   r�   r°   r·   rÇ   rÏ   râ   r  r  r&  r<  rC  rh  rp  rm  r{  r©  rÊ  ÚFLAX_WAV2VEC2_MODEL_DOCSTRINGrÍ  r×  ÚFLAX_WAV2VEC2_FOR_CTC_DOCSTRINGrÙ  ré  Ú'FLAX_WAV2VEC2_FOR_PRETRAINING_DOCSTRINGÚ__all__r(   r)   r*   ú<module>r     sš  ðñ å ß )Ñ )ã Ý Û 
Ý Û ß >Ñ >Ý >ß ;Ý ç L÷ó ÷ gÓ fÝ 2ð 
ˆ×	Ñ	˜HÓ	%€ð ‡�×Ñô4 +ó 4ó ð4ð: ‡�×Ñô4 {ó 4ó ð4ðL ,0ØñFØ��c�‰?ðFàðFð ðFð ˜RŸZ™ZÑ(ð	Fð
 ðFð ‡Z�ZóFñR$¨Uð $À3ð $ÐX`Ðac×akÑakÑXló $ðB$Ð ðN Ð ôF R§Y¡Yô ô8!˜RŸY™Yô !ôH¨"¯)©)ô ô,˜rŸy™yô ô0 §¡ô ô"1 B§I¡Iô 1ô(X)˜BŸI™Iô X)ôv˜bŸi™iô ôD"¨b¯i©iô "ôJ-
¸¿	¹	ô -
ô`4
¨¯©ô 4
ônK'¨¯	©	ô K'ô\˜"Ÿ)™)ô ô:˜rŸy™yô ô*¨"¯)©)ô ô"ZdÐ"5ô Zdôzm˜Ÿ™ô mñ` ØhØóô&Ð3ó &ó	ð&ð!Ð ñ< ØØÐ =Ñ=ôñ !ØÐ#>È^õô
D˜rŸy™yô DñN ØlØóô,Ð4ó ,ó	ð,ð!#Ð ñF ØØÐ ?Ñ?ôñ !Ð!3ÐASÐbpÕ qô` r§y¡yô `ñF ÐQÐSkÓlô5
Ð!<ó 5
ó mð5
ðp*+Ð 'ñX ØØÐ GÑGôñ !ØÐ,LÐ[iõò
 s�r)   