Ë
    T^(hœc  ã                   ó¨  — d dl mZmZ d dlmZ d dlZd dlm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mZmZ ddlmZmZmZmZ dd	lmZmZ d
dlm Z  dZ!dZ" G d„ dejF                  «      Z$ G d„ dejF                  «      Z% G d„ dejF                  «      Z& G d„ dejF                  «      Z' G d„ dejF                  «      Z( G d„ dejF                  «      Z) G d„ dejF                  «      Z* G d„ dejF                  «      Z+ G d„ dejF                  «      Z, G d „ d!ejF                  «      Z- G d"„ d#ejF                  «      Z. G d$„ d%e«      Z/ G d&„ d'ejF                  «      Z0 ed(e!«       G d)„ d*e/«      «       Z1d+Z2 ee1e2«        ee1ee ¬,«        G d-„ d.ejF                  «      Z3 ed/e!«       G d0„ d1e/«      «       Z4d2Z5 ee4e5«        ee4ee ¬,«       g d3¢Z6y)4é    )ÚOptionalÚTupleN)Ú
FrozenDictÚfreezeÚunfreeze)Údot_product_attention_weights)Úflatten_dictÚunflatten_dicté   )ÚFlaxBaseModelOutputÚFlaxBaseModelOutputWithPoolingÚFlaxSequenceClassifierOutput)ÚACT2FNÚFlaxPreTrainedModelÚ append_replace_return_docstringsÚoverwrite_call_docstring)Úadd_start_docstringsÚ%add_start_docstrings_to_model_forwardé   )Ú	ViTConfigaû  

    This model inherits from [`FlaxPreTrainedModel`]. Check the superclass documentation for the generic methods the
    library implements for all its model (such as downloading, saving and converting weights from PyTorch models)

    This model is also a
    [flax.linen.Module](https://flax.readthedocs.io/en/latest/api_reference/flax.linen/module.html) subclass. Use it as
    a regular Flax linen 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 ([`ViTConfig`]): 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:
        pixel_values (`numpy.ndarray` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See [`ViTImageProcessor.__call__`]
            for details.

        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                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)ÚFlaxViTPatchEmbeddingsÚconfigÚdtypec                 óº  — | j                   j                  }| j                   j                  }||z  ||z  z  }|| _        | j                   j                  | _        t        j                  | j                   j                  ||f||fd| j                  t        j
                  j                  j                  | j                   j                  dz  dd«      ¬«      | _        y )NÚVALIDé   Úfan_inÚtruncated_normal)Úkernel_sizeÚstridesÚpaddingr   Úkernel_init)r   Ú
image_sizeÚ
patch_sizeÚnum_patchesÚnum_channelsÚnnÚConvÚhidden_sizer   ÚjaxÚinitializersÚvariance_scalingÚinitializer_rangeÚ
projection)Úselfr$   r%   r&   s       úg/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/vit/modeling_flax_vit.pyÚsetupzFlaxViTPatchEmbeddings.setup\   s´   € Ø—[‘[×+Ñ+ˆ
Ø—[‘[×+Ñ+ˆ
Ø! ZÑ/°JÀ*Ñ4LÑMˆØ&ˆÔØ ŸK™K×4Ñ4ˆÔÜŸ'™'Ø�K‰K×#Ñ#Ø# ZÐ0Ø Ð,ØØ—*‘*ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°(Ð<Nóô	
ˆ�ó    c                 óÊ   — |j                   d   }|| j                  k7  rt        d«      ‚| j                  |«      }|j                   \  }}}}t	        j
                  ||d|f«      S )NéÿÿÿÿzeMake sure that the channel dimension of the pixel values match with the one set in the configuration.)Úshaper'   Ú
ValueErrorr/   ÚjnpÚreshape)r0   Úpixel_valuesr'   Ú
embeddingsÚ
batch_sizeÚ_Úchannelss          r1   Ú__call__zFlaxViTPatchEmbeddings.__call__m   sl   € Ø#×)Ñ)¨"Ñ-ˆØ˜4×,Ñ,Ò,ÜØwóð ð —_‘_ \Ó2ˆ
Ø%/×%5Ñ%5Ñ"ˆ
�A�q˜(Ü�{‰{˜:¨
°B¸Ð'AÓBÐBr3   N©
Ú__name__Ú
__module__Ú__qualname__r   Ú__annotations__r8   Úfloat32r   r2   r?   © r3   r1   r   r   X   s%   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ó"Cr3   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)ÚFlaxViTEmbeddingsz7Construct the CLS token, position and patch embeddings.r   r   c                 óœ  — | j                  dt        j                  j                  j	                  | j
                  j                  dz  dd«      dd| j
                  j                  f«      | _        t        | j
                  | j                  ¬«      | _        | j                  j                  }| j                  dt        j                  j                  j	                  | j
                  j                  dz  dd«      d|dz   | j
                  j                  f«      | _        t        j                  | j
                  j                  ¬«      | _        y )	NÚ	cls_tokenr   r   r   r   ©r   Úposition_embeddings©Úrate)Úparamr+   r(   r,   r-   r   r.   r*   rJ   r   r   Úpatch_embeddingsr&   rL   ÚDropoutÚhidden_dropout_probÚdropout)r0   r&   s     r1   r2   zFlaxViTEmbeddings.setup~   s÷   € ØŸ™ØÜ�F‰F×Ñ×0Ñ0°·±×1NÑ1NÐPQÑ1QÐS[Ð]oÓpØ��4—;‘;×*Ñ*Ð+ó
ˆŒô
 !7°t·{±{È$Ï*É*Ô UˆÔØ×+Ñ+×7Ñ7ˆØ#'§:¡:Ø!Ü�F‰F×Ñ×0Ñ0°·±×1NÑ1NÐPQÑ1QÐS[Ð]oÓpØ�˜a‘ §¡×!8Ñ!8Ð9ó$
ˆÔ ô
 —z‘z t§{¡{×'FÑ'FÔGˆ�r3   c                 ó*  — |j                   d   }| j                  |«      }t        j                  | j                  |d| j
                  j                  f«      }t        j                  ||fd¬«      }|| j                  z   }| j                  ||¬«      }|S )Nr   r   )Úaxis©Údeterministic)
r6   rP   r8   Úbroadcast_torJ   r   r*   ÚconcatenaterL   rS   )r0   r:   rW   r<   r;   Ú
cls_tokenss         r1   r?   zFlaxViTEmbeddings.__call__�   s†   € Ø!×'Ñ'¨Ñ*ˆ
à×*Ñ*¨<Ó8ˆ
ä×%Ñ% d§n¡n°zÀ1ÀdÇkÁk×F]ÑF]Ð6^Ó_ˆ
Ü—_‘_ j°*Ð%=ÀAÔFˆ
Ø $×":Ñ":Ñ:ˆ
Ø—\‘\ *¸M�\ÓJˆ
ØÐr3   N©T)rA   rB   rC   Ú__doc__r   rD   r8   rE   r   r2   r?   rF   r3   r1   rH   rH   x   s(   … ÙAàÓØ—{‘{€Eˆ3�9‰9Ó"òHô	r3   rH   c                   óf   — 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
y)	ÚFlaxViTSelfAttentionr   r   c           	      óà  — | j                   j                  | j                   j                  z  dk7  rt        d«      ‚t	        j
                  | j                   j                  | j                  t        j                  j                  j                  | j                   j                  dz  dd¬«      | j                   j                  ¬«      | _        t	        j
                  | j                   j                  | j                  t        j                  j                  j                  | j                   j                  dz  dd¬«      | j                   j                  ¬«      | _        t	        j
                  | j                   j                  | j                  t        j                  j                  j                  | j                   j                  dz  dd¬«      | j                   j                  ¬«      | _        y )Nr   z‡`config.hidden_size`: {self.config.hidden_size} has to be a multiple of `config.num_attention_heads`: {self.config.num_attention_heads}r   r   r   )ÚmodeÚdistribution)r   r#   Úuse_bias)r   r*   Únum_attention_headsr7   r(   ÚDenser   r+   r,   r-   r.   Úqkv_biasÚqueryÚkeyÚvalue©r0   s    r1   r2   zFlaxViTSelfAttention.setup�   sp  € Ø�;‰;×"Ñ" T§[¡[×%DÑ%DÑDÈÒIÜð5óð ô
 —X‘XØ�K‰K×#Ñ#Ø—*‘*ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°xÐN`ð =ó ð —[‘[×)Ñ)ô
ˆŒ
ô —8‘8Ø�K‰K×#Ñ#Ø—*‘*ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°xÐN`ð =ó ð —[‘[×)Ñ)ô
ˆŒô —X‘XØ�K‰K×#Ñ#Ø—*‘*ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°xÐN`ð =ó ð —[‘[×)Ñ)ô
ˆ�
r3   rW   Úoutput_attentionsc           
      óH  — | j                   j                  | j                   j                  z  }| j                  |«      j	                  |j
                  d d | j                   j                  |fz   «      }| j                  |«      j	                  |j
                  d d | j                   j                  |fz   «      }| j                  |«      j	                  |j
                  d d | j                   j                  |fz   «      }d }|s*| j                   j                  dkD  r| j                  d«      }t        |||| j                   j                  d|| j                  d ¬«      }	t        j                  d|	|«      }
|
j	                  |
j
                  d d dz   «      }
|r|
|	f}|S |
f}|S )Nr   g        rS   T)Údropout_rngÚdropout_rateÚbroadcast_dropoutrW   r   Ú	precisionz...hqk,...khd->...qhd)r5   )r   r*   rc   rf   r9   r6   rh   rg   Úattention_probs_dropout_probÚmake_rngr   r   r8   Úeinsum)r0   Úhidden_statesrW   rj   Úhead_dimÚquery_statesÚvalue_statesÚ
key_statesrl   Úattn_weightsÚattn_outputÚoutputss               r1   r?   zFlaxViTSelfAttention.__call__½   sŒ  € Ø—;‘;×*Ñ*¨d¯k©k×.MÑ.MÑMˆà—z‘z -Ó0×8Ñ8Ø×Ñ  Ð# t§{¡{×'FÑ'FÈÐ&QÑQó
ˆð —z‘z -Ó0×8Ñ8Ø×Ñ  Ð# t§{¡{×'FÑ'FÈÐ&QÑQó
ˆð —X‘X˜mÓ,×4Ñ4Ø×Ñ  Ð# t§{¡{×'FÑ'FÈÐ&QÑQó
ˆ
ð ˆÙ §¡×!IÑ!IÈCÒ!OØŸ-™-¨	Ó2ˆKä4ØØØ#ØŸ™×AÑAØ"Ø'Ø—*‘*Øô	
ˆô —j‘jÐ!8¸,ÈÓUˆØ!×)Ñ)¨+×*;Ñ*;¸B¸QÐ*?À%Ñ*GÓHˆá1B�; Ð-ˆØˆð JUÈˆØˆr3   N©TF©rA   rB   rC   r   rD   r8   rE   r   r2   Úboolr?   rF   r3   r1   r^   r^   ™   s4   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ñ@ °Tð  ÐUYô  r3   r^   c                   ób   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdde	fd„Z
y)ÚFlaxViTSelfOutputr   r   c                 óX  — t        j                  | j                  j                  t        j                   j
                  j                  | j                  j                  dz  dd«      | j                  ¬«      | _	        t        j                  | j                  j                  ¬«      | _        y ©Nr   r   r   ©r#   r   rM   ©r(   rd   r   r*   r+   r,   r-   r.   r   ÚdenserQ   rR   rS   ri   s    r1   r2   zFlaxViTSelfOutput.setupä   ós   € Ü—X‘XØ�K‰K×#Ñ#ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°(Ð<Nóð —*‘*ô
ˆŒ
ô —z‘z t§{¡{×'FÑ'FÔGˆ�r3   rW   c                 óN   — | j                  |«      }| j                  ||¬«      }|S ©NrV   ©r„   rS   )r0   rs   Úinput_tensorrW   s       r1   r?   zFlaxViTSelfOutput.__call__î   s(   € ØŸ
™
 =Ó1ˆØŸ™ ]À-˜ÓPˆØÐr3   Nr[   r|   rF   r3   r1   r   r   à   s,   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òHñÀ4ô r3   r   c                   ób   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdde	fd„Z
y)ÚFlaxViTAttentionr   r   c                 óœ   — t        | j                  | j                  ¬«      | _        t	        | j                  | j                  ¬«      | _        y ©NrK   )r^   r   r   Ú	attentionr   Úoutputri   s    r1   r2   zFlaxViTAttention.setupø   s.   € Ü-¨d¯k©kÀÇÁÔLˆŒÜ'¨¯©¸4¿:¹:ÔFˆ�r3   rj   c                 ó|   — | j                  |||¬«      }|d   }| j                  |||¬«      }|f}|r	||d   fz  }|S ©N©rW   rj   r   rV   r   )rŽ   r�   )r0   rs   rW   rj   Úattn_outputsry   rz   s          r1   r?   zFlaxViTAttention.__call__ü   sU   € Ø—~‘~ mÀ=Ðdu�~ÓvˆØ" 1‘oˆØŸ™ K°Èm˜Ó\ˆà Ð"ˆáØ˜ Q™Ð)Ñ)ˆGàˆr3   Nr{   r|   rF   r3   r1   r‹   r‹   ô   s,   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òGñ
ÈTô 
r3   r‹   c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)ÚFlaxViTIntermediater   r   c                 ó>  — t        j                  | j                  j                  t        j                   j
                  j                  | j                  j                  dz  dd«      | j                  ¬«      | _	        t        | j                  j                     | _        y ©Nr   r   r   r‚   )r(   rd   r   Úintermediate_sizer+   r,   r-   r.   r   r„   r   Ú
hidden_actÚ
activationri   s    r1   r2   zFlaxViTIntermediate.setup  so   € Ü—X‘XØ�K‰K×)Ñ)ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°(Ð<Nóð —*‘*ô
ˆŒ
ô ! §¡×!7Ñ!7Ñ8ˆ�r3   c                 óJ   — | j                  |«      }| j                  |«      }|S ©N©r„   rš   )r0   rs   s     r1   r?   zFlaxViTIntermediate.__call__  s$   € ØŸ
™
 =Ó1ˆØŸ™¨Ó6ˆØÐr3   Nr@   rF   r3   r1   r•   r•   	  s$   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò9ór3   r•   c                   ób   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdde	fd„Z
y)ÚFlaxViTOutputr   r   c                 óX  — t        j                  | j                  j                  t        j                   j
                  j                  | j                  j                  dz  dd«      | j                  ¬«      | _	        t        j                  | j                  j                  ¬«      | _        y r�   rƒ   ri   s    r1   r2   zFlaxViTOutput.setup!  r…   r3   rW   c                 óX   — | j                  |«      }| j                  ||¬«      }||z   }|S r‡   rˆ   )r0   rs   Úattention_outputrW   s       r1   r?   zFlaxViTOutput.__call__+  s3   € ØŸ
™
 =Ó1ˆØŸ™ ]À-˜ÓPˆØ%Ð(8Ñ8ˆØÐr3   Nr[   r|   rF   r3   r1   rŸ   rŸ     s,   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òHñÀtô r3   rŸ   c                   óf   — 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
y)	ÚFlaxViTLayerr   r   c                 óÐ  — t        | j                  | j                  ¬«      | _        t	        | j                  | j                  ¬«      | _        t        | j                  | j                  ¬«      | _        t        j                  | j                  j                  | j                  ¬«      | _        t        j                  | j                  j                  | j                  ¬«      | _        y ©NrK   )Úepsilonr   )r‹   r   r   rŽ   r•   ÚintermediaterŸ   r�   r(   Ú	LayerNormÚlayer_norm_epsÚlayernorm_beforeÚlayernorm_afterri   s    r1   r2   zFlaxViTLayer.setup6  s�   € Ü)¨$¯+©+¸T¿Z¹ZÔHˆŒÜ/°·±À4Ç:Á:ÔNˆÔÜ# D§K¡K°t·z±zÔBˆŒÜ "§¡°T·[±[×5OÑ5OÐW[×WaÑWaÔ bˆÔÜ!Ÿ|™|°D·K±K×4NÑ4NÐVZ×V`ÑV`ÔaˆÕr3   rW   rj   c                 óè   — | j                  | j                  |«      ||¬«      }|d   }||z   }| j                  |«      }| j                  |«      }| j	                  |||¬«      }|f}|r	||d   fz  }|S r‘   )rŽ   r«   r¬   r¨   r�   )r0   rs   rW   rj   Úattention_outputsr¢   Úlayer_outputrz   s           r1   r?   zFlaxViTLayer.__call__=  s    € Ø ŸN™NØ×!Ñ! -Ó0Ø'Ø/ð +ó 
Ðð -¨QÑ/Ðð ,¨mÑ;Ðð ×+Ñ+Ð,<Ó=ˆà×)Ñ)¨,Ó7ˆØŸ™ MÐ3CÐS`˜Óaˆà Ð"ˆáØÐ)¨!Ñ,Ð.Ñ.ˆGØˆr3   Nr{   r|   rF   r3   r1   r¤   r¤   2  s4   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òbñ°Tð ÐUYô r3   r¤   c            	       óv   — 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	)ÚFlaxViTLayerCollectionr   r   c           	      óÄ   — t        | j                  j                  «      D �cg c]-  }t        | j                  t	        |«      | j
                  ¬«      ‘Œ/ c}| _        y c c}w )N)Únamer   )Úranger   Únum_hidden_layersr¤   Ústrr   Úlayers)r0   Úis     r1   r2   zFlaxViTLayerCollection.setupZ  sE   € äNSÐTX×T_ÑT_×TqÑTqÓNrö
ØIJŒL˜Ÿ™¬3¨q«6¸¿¹ÖDò
ˆ�ùò 
s   ¢2ArW   rj   Ú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 )NrF   r’   r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrœ   rF   )Ú.0Úvs     r1   ú	<genexpr>z2FlaxViTLayerCollection.__call__.<locals>.<genexpr>z  s   è ø€ Ò=˜q¨q©}œÑ=ùs   ‚Š)Úlast_hidden_staters   Ú
attentions)Ú	enumerater·   Útupler   )r0   rs   rW   rj   r¹   rº   Úall_attentionsÚall_hidden_statesr¸   ÚlayerÚlayer_outputsrz   s               r1   r?   zFlaxViTLayerCollection.__call___  s·   € ñ  1™°dˆÙ"6™B¸DÐä! $§+¡+Ó.ò 		6‰HˆAˆuÙ#Ø! mÐ%5Ñ5Ð!á! -¸}Ð`qÔrˆMà)¨!Ñ,ˆMâ Ø =°Ñ#3Ð"5Ñ5‘ð		6ñ  Ø -Ð!1Ñ1Ðà Ð"ˆÙÜÑ= GÔ=Ó=Ð=ä"Ø+Ð;LÐYgô
ð 	
r3   N©TFFTr|   rF   r3   r1   r±   r±   V  sZ   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð #Ø"'Ø%*Ø ñ
ð ð
ð  ð	
ð
 #ð
ð ô
r3   r±   c            	       óv   — 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	)ÚFlaxViTEncoderr   r   c                 óP   — t        | j                  | j                  ¬«      | _        y r�   )r±   r   r   rÆ   ri   s    r1   r2   zFlaxViTEncoder.setup…  s   € Ü+¨D¯K©K¸t¿z¹zÔJˆ�
r3   rW   rj   r¹   rº   c                 ó.   — | j                  |||||¬«      S )N©rW   rj   r¹   rº   )rÆ   )r0   rs   rW   rj   r¹   rº   s         r1   r?   zFlaxViTEncoder.__call__ˆ  s)   € ð �z‰zØØ'Ø/Ø!5Ø#ð ó 
ð 	
r3   NrÈ   r|   rF   r3   r1   rÊ   rÊ   �  s[   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òKð #Ø"'Ø%*Ø ñ
ð ð
ð  ð	
ð
 #ð
ð ô
r3   rÊ   c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)ÚFlaxViTPoolerr   r   c                 ó>  — t        j                  | j                  j                  t        j                   j
                  j                  | j                  j                  dz  dd«      | j                  ¬«      | _	        t        | j                  j                     | _        y r—   )r(   rd   r   Úpooler_output_sizer+   r,   r-   r.   r   r„   r   Ú
pooler_actrš   ri   s    r1   r2   zFlaxViTPooler.setup�  so   € Ü—X‘XØ�K‰K×*Ñ*ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°(Ð<Nóð —*‘*ô
ˆŒ
ô ! §¡×!7Ñ!7Ñ8ˆ�r3   c                 óX   — |d d …df   }| j                  |«      }| j                  |«      S )Nr   r�   )r0   rs   Úcls_hidden_states      r1   r?   zFlaxViTPooler.__call__§  s1   € Ø(ª¨A¨Ñ.ÐØŸ:™:Ð&6Ó7ÐØ�‰Ð/Ó0Ð0r3   Nr@   rF   r3   r1   rÏ   rÏ   ™  s$   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò9ó1r3   rÏ   c                   ót  ‡ — e Zd ZU dZeZdZdZdZe	j                  ed<   ddej                  df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j5                  d«      «      	 	 	 	 	 	 ddedej&                  j(                  dedee   dee   dee   fd„«       Zˆ xZS )ÚFlaxViTPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úvitr:   NÚmodule_classr   Tr   Úseedr   Ú_do_initc                 ó¦   •—  | j                   d||dœ|¤Ž}|€$d|j                  |j                  |j                  f}t        ‰| �  ||||||¬«       y )N)r   r   r   )Úinput_shaperÙ   r   rÚ   rF   )rØ   r$   r'   ÚsuperÚ__init__)	r0   r   rÜ   rÙ   r   rÚ   ÚkwargsÚmoduleÚ	__class__s	           €r1   rÞ   zFlaxViTPreTrainedModel.__init__¸  sc   ø€ ð #�×"Ñ"ÐH¨&¸ÑHÀÑHˆØÐØ˜f×/Ñ/°×1BÑ1BÀF×DWÑDWÐXˆKÜ‰Ñ˜ °[ÀtÐSXÐckÐÕlr3   ÚrngrÜ   ÚparamsÚreturnc                 ó¤  — t        j                  || j                  ¬«      }t        j                  j                  |«      \  }}||dœ}| j                  j                  ||d¬«      d   }|�dt        t        |«      «      }t        t        |«      «      }| j                  D ]
  }	||	   ||	<   Œ t        «       | _
        t        t        |«      «      S |S )NrK   )rã   rS   F)rº   rã   )r8   Úzerosr   r+   ÚrandomÚsplitrà   Úinitr	   r   Ú_missing_keysÚsetr   r
   )
r0   râ   rÜ   rã   r:   Ú
params_rngrl   ÚrngsÚrandom_paramsÚmissing_keys
             r1   Úinit_weightsz#FlaxViTPreTrainedModel.init_weightsÆ  sÄ   € ä—y‘y °D·J±JÔ?ˆä"%§*¡*×"2Ñ"2°3Ó"7Ñˆ
�KØ$°Ñ=ˆàŸ™×(Ñ(¨¨|ÈÐ(ÓOÐPXÑYˆàÐÜ(¬°-Ó)@ÓAˆMÜ!¤(¨6Ó"2Ó3ˆFØ#×1Ñ1ò A�Ø&3°KÑ&@��{Ò#ðAä!$£ˆDÔÜœ.¨Ó0Ó1Ð1à Ð r3   zbatch_size, sequence_lengthrl   Útrainrj   r¹   rº   c           	      óˆ  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }t	        j
                  |d«      }i }|�||d<   | j                  j                  d|xs | j                  it	        j                  |t        j                  ¬«      | ||||¬«      S )N)r   r   r   r   rS   rã   rK   )rí   )r   rj   r¹   rº   r8   Ú	transposerà   Úapplyrã   ÚarrayrE   )	r0   r:   rã   rl   rñ   rj   r¹   rº   rí   s	            r1   r?   zFlaxViTPreTrainedModel.__call__Ù  sÆ   € ð 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆä—}‘} \°<Ó@ˆàˆØÐ"Ø)ˆD�‰Oà�{‰{× Ñ Ø�vÒ, §¡Ð-Ü�I‰I�l¬#¯+©+Ô6ØˆIØØ ØØð !ó 
ð 	
r3   rœ   )NNFNNN)rA   rB   rC   r\   r   Úconfig_classÚbase_model_prefixÚmain_input_namerØ   r(   ÚModulerD   r8   rE   Úintr   r}   rÞ   r+   rç   ÚPRNGKeyr   r   rð   r   ÚVIT_INPUTS_DOCSTRINGÚformatÚdictr   r?   Ú__classcell__)rá   s   @r1   rÖ   rÖ   ­  s(  ø… ñð
 €LØÐØ$€OØ"€L�"—)‘)Ó"ð
 ØØŸ;™;Øñmàðmð ð	mð
 �y‰yðmð õmñ! §
¡
× 2Ñ 2ð !Àð !ÐPZð !Ðfpó !ñ& +Ð+?×+FÑ+FÐGdÓ+eÓfð Ø*.ØØ,0Ø/3Ø&*ñ
ð ð
ð —Z‘Z×'Ñ'ð	
ð
 ð
ð $ D™>ð
ð ' t™nð
ð ˜d‘^ò
ó gô
r3   rÖ   c            	       ó„   — e Zd ZU eed<   ej                  Zej                  ed<   dZe	ed<   d„ Z
	 	 	 	 dde	de	de	d	e	fd
„Zy)ÚFlaxViTModuler   r   TÚadd_pooling_layerc                 ó„  — t        | j                  | j                  ¬«      | _        t	        | j                  | j                  ¬«      | _        t        j                  | j                  j                  | j                  ¬«      | _	        | j                  r't        | j                  | j                  ¬«      | _        y d | _        y r¦   )rH   r   r   r;   rÊ   Úencoderr(   r©   rª   Ú	layernormr  rÏ   Úpoolerri   s    r1   r2   zFlaxViTModule.setup   sv   € Ü+¨D¯K©K¸t¿z¹zÔJˆŒÜ% d§k¡k¸¿¹ÔDˆŒÜŸ™¨d¯k©k×.HÑ.HÐPT×PZÑPZÔ[ˆŒØFJ×F\ÒF\”m D§K¡K°t·z±zÔBˆ�Ðbfˆ�r3   rW   rj   r¹   rº   c                 ó2  — | j                  ||¬«      }| j                  |||||¬«      }|d   }| j                  |«      }| j                  r| j	                  |«      nd }|s|€	|f|dd  z   S ||f|dd  z   S t        |||j                  |j                  ¬«      S )NrV   rÍ   r   r   )rÀ   Úpooler_outputrs   rÁ   )r;   r  r  r  r  r   rs   rÁ   )	r0   r:   rW   rj   r¹   rº   rs   rz   Úpooleds	            r1   r?   zFlaxViTModule.__call__  sÀ   € ð Ÿ™¨ÀM˜ÓRˆà—,‘,ØØ'Ø/Ø!5Ø#ð ó 
ˆð   ™
ˆØŸ™ }Ó5ˆØ/3×/EÒ/E�—‘˜]Ô+È4ˆáàˆ~Ø%Ð'¨'°!°"¨+Ñ5Ð5Ø! 6Ð*¨W°Q°R¨[Ñ8Ð8ä-Ø+Ø Ø!×/Ñ/Ø×)Ñ)ô	
ð 	
r3   NrÈ   )rA   rB   rC   r   rD   r8   rE   r   r  r}   r2   r?   rF   r3   r1   r  r  û  sf   … ØÓØ—{‘{€Eˆ3�9‰9Ó"Ø"Ð�tÓ"ògð #Ø"'Ø%*Ø ñ 
ð ð 
ð  ð	 
ð
 #ð 
ð ô 
r3   r  z]The bare ViT Model transformer outputting raw hidden-states without any specific head on top.c                   ó   — e Zd ZeZy)ÚFlaxViTModelN)rA   rB   rC   r  rØ   rF   r3   r1   r  r  )  s	   „ ð
 !�Lr3   r  a†  
    Returns:

    Examples:

    ```python
    >>> from transformers import AutoImageProcessor, FlaxViTModel
    >>> from PIL import Image
    >>> import requests

    >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
    >>> image = Image.open(requests.get(url, stream=True).raw)

    >>> image_processor = AutoImageProcessor.from_pretrained("google/vit-base-patch16-224-in21k")
    >>> model = FlaxViTModel.from_pretrained("google/vit-base-patch16-224-in21k")

    >>> inputs = image_processor(images=image, return_tensors="np")
    >>> outputs = model(**inputs)
    >>> last_hidden_states = outputs.last_hidden_state
    ```
)Úoutput_typerö   c                   ól   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 	 dde	fd„Z
y)Ú#FlaxViTForImageClassificationModuler   r   c           	      óH  — t        | j                  | j                  d¬«      | _        t	        j
                  | j                  j                  | j                  t        j                  j                  j                  | j                  j                  dz  dd«      ¬«      | _        y )NF)r   r   r  r   r   r   )r   r#   )r  r   r   r×   r(   rd   Ú
num_labelsr+   r,   r-   r.   Ú
classifierri   s    r1   r2   z)FlaxViTForImageClassificationModule.setupO  sn   € Ü ¨¯©¸4¿:¹:ÐY^Ô_ˆŒÜŸ(™(Ø�K‰K×"Ñ"Ø—*‘*ÜŸ™×+Ñ+×<Ñ<Ø—‘×-Ñ-¨qÑ0°(Ð<Nóô
ˆ�r3   NrW   c                 ó   — |�|n| j                   j                  }| j                  |||||¬«      }|d   }| j                  |d d …dd d …f   «      }|s|f|dd  z   }	|	S t	        ||j
                  |j                  ¬«      S )NrÍ   r   r   )Úlogitsrs   rÁ   )r   Úuse_return_dictr×   r  r   rs   rÁ   )
r0   r:   rW   rj   r¹   rº   rz   rs   r  r�   s
             r1   r?   z,FlaxViTForImageClassificationModule.__call__Y  sœ   € ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—(‘(ØØ'Ø/Ø!5Ø#ð ó 
ˆð   ™
ˆØ—‘ ªq°!²Q¨wÑ!7Ó8ˆáØ�Y ¨¨ Ñ,ˆFØˆMä+ØØ!×/Ñ/Ø×)Ñ)ô
ð 	
r3   )NTNNNr|   rF   r3   r1   r  r  K  s?   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð Ø"ØØ!Øñ
ð ô
r3   r  z¤
    ViT Model transformer with an image classification head on top (a linear layer on top of the final hidden state of
    the [CLS] token) e.g. for ImageNet.
    c                   ó   — e Zd ZeZy)ÚFlaxViTForImageClassificationN)rA   rB   rC   r  rØ   rF   r3   r1   r  r  y  s	   „ ð 7�Lr3   r  ag  
    Returns:

    Example:

    ```python
    >>> from transformers import AutoImageProcessor, FlaxViTForImageClassification
    >>> from PIL import Image
    >>> import jax
    >>> import requests

    >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
    >>> image = Image.open(requests.get(url, stream=True).raw)

    >>> image_processor = AutoImageProcessor.from_pretrained("google/vit-base-patch16-224")
    >>> model = FlaxViTForImageClassification.from_pretrained("google/vit-base-patch16-224")

    >>> inputs = image_processor(images=image, return_tensors="np")
    >>> outputs = model(**inputs)
    >>> logits = outputs.logits

    >>> # model predicts one of the 1000 ImageNet classes
    >>> predicted_class_idx = jax.numpy.argmax(logits, axis=-1)
    >>> print("Predicted class:", model.config.id2label[predicted_class_idx.item()])
    ```
)r  r  rÖ   )7Útypingr   r   Ú
flax.linenÚlinenr(   r+   Ú	jax.numpyÚnumpyr8   Úflax.core.frozen_dictr   r   r   Úflax.linen.attentionr   Úflax.traverse_utilr	   r
   Úmodeling_flax_outputsr   r   r   Úmodeling_flax_utilsr   r   r   r   Úutilsr   r   Úconfiguration_vitr   ÚVIT_START_DOCSTRINGrü   rù   r   rH   r^   r   r‹   r•   rŸ   r¤   r±   rÊ   rÏ   rÖ   r  r  ÚFLAX_VISION_MODEL_DOCSTRINGr  r  ÚFLAX_VISION_CLASSIF_DOCSTRINGÚ__all__rF   r3   r1   ú<module>r'     sØ  ð÷  #å Û 
Ý ß >Ñ >Ý >ß ;ç vÑ v÷ó ÷ QÝ (ð!Ð ðFÐ ô"C˜RŸY™Yô Cô@˜Ÿ	™	ô ôBD˜2Ÿ9™9ô DôN˜Ÿ	™	ô ô(�r—y‘yô ô*˜"Ÿ)™)ô ô(�B—I‘Iô ô*!�2—9‘9ô !ôH(
˜RŸY™Yô (
ôV
�R—Y‘Yô 
ô01�B—I‘Iô 1ô(K
Ð0ô K
ô\+
�B—I‘Iô +
ñ\ ØcØóô!Ð)ó !ó	ð!ðÐ ñ, ˜Ð'BÔ CÙ   Ð;YÐhqÕ rô+
¨"¯)©)ô +
ñ\ ðð óô7Ð$:ó 7óð7ð!Ð ñ6 Ð6Ð8UÔ VÙ  Ø!Ð/KÐZcõò
 V�r3   