Ë
    T^(h€`  ã                   óü  — d dl mZ 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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jF                  «      Z/ G d&„ d'e«      Z0 G d(„ d)ejF                  «      Z1 ed*e!«       G d+„ d,e0«      «       Z2d-Z3 ee2e3«        ee2ee ¬.«        G d/„ d0ejF                  «      Z4 G d1„ d2ejF                  «      Z5 ed3e!«       G d4„ d5e0«      «       Z6d6Z7 ee6e7«        ee6ee ¬.«       g d7¢Z8y)8é    )Úpartial)ÚOptionalÚTupleN)Ú
FrozenDictÚfreezeÚunfreeze)Úflatten_dictÚunflatten_dicté   )Ú"FlaxBaseModelOutputWithNoAttentionÚ,FlaxBaseModelOutputWithPoolingAndNoAttentionÚ(FlaxImageClassifierOutputWithNoAttention)ÚACT2FNÚFlaxPreTrainedModelÚ append_replace_return_docstringsÚoverwrite_call_docstring)Úadd_start_docstringsÚ%add_start_docstrings_to_model_forwardé   )ÚResNetConfigaþ  

    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 ([`ResNetConfig`]): 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`].
aA  
    Args:
        pixel_values (`jax.numpy.float32` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See
            [`AutoImageProcessor.__call__`] for details.
        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                   ó4   — e Zd ZdZej
                  d„ «       Zy)ÚIdentityzIdentity function.c                 ó   — |S ©N© )ÚselfÚxÚkwargss      úm/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/resnet/modeling_flax_resnet.pyÚ__call__zIdentity.__call__\   s   € àˆó    N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚnnÚcompactr    r   r!   r   r   r   Y   s   „ Ùà‡Z�Zñó ñr!   r   c                   óÂ   — e Zd ZU eed<   dZeed<   dZeed<   dZee	   ed<   e
j                  Ze
j                  ed<   d	„ Zdd
e
j                  dede
j                  fd„Zy)ÚFlaxResNetConvLayerÚout_channelsr   Úkernel_sizer   ÚstrideÚreluÚ
activationÚdtypec                 óÔ  — t        j                  | j                  | j                  | j                  f| j                  | j                  dz  | j
                  dt         j                  j                  ddd| j
                  ¬«      ¬«      | _        t        j                  dd	| j
                  ¬
«      | _
        | j                  �t        | j                     | _        y t        «       | _        y )Né   Fç       @Úfan_outÚnormal)ÚmodeÚdistributionr/   )r+   ÚstridesÚpaddingr/   Úuse_biasÚkernel_initçÍÌÌÌÌÌì?çñhãˆµøä>©ÚmomentumÚepsilonr/   )r&   ÚConvr*   r+   r,   r/   ÚinitializersÚvariance_scalingÚconvolutionÚ	BatchNormÚnormalizationr.   r   r   Úactivation_func©r   s    r   ÚsetupzFlaxResNetConvLayer.setuph   s²   € ÜŸ7™7Ø×ÑØ×)Ñ)¨4×+;Ñ+;Ð<Ø—K‘KØ×$Ñ$¨Ñ)Ø—*‘*ØÜŸ™×8Ñ8¸À9Ð[cÐko×kuÑkuÐ8Óvô
ˆÔô  Ÿ\™\°3ÀÈTÏZÉZÔXˆÔØ:>¿/¹/Ð:Uœv d§o¡oÑ6ˆÕÔ[cÓ[eˆÕr!   r   ÚdeterministicÚreturnc                 óp   — | j                  |«      }| j                  ||¬«      }| j                  |«      }|S ©N)Úuse_running_average)rC   rE   rF   ©r   r   rI   Úhidden_states       r   r    zFlaxResNetConvLayer.__call__u   s=   € Ø×'Ñ'¨Ó*ˆØ×)Ñ)¨,ÈMÐ)ÓZˆØ×+Ñ+¨LÓ9ˆØÐr!   N©T)r"   r#   r$   ÚintÚ__annotations__r+   r,   r.   r   ÚstrÚjnpÚfloat32r/   rH   ÚndarrayÚboolr    r   r!   r   r)   r)   a   sc   … ØÓØ€K�ÓØ€FˆCƒOØ &€J�˜‘Ó&Ø—{‘{€Eˆ3�9‰9Ó"òfñ˜#Ÿ+™+ð °dð ÀcÇkÁkô r!   r)   c                   ó–   — e Zd ZU dZeed<   ej                  Zej                  ed<   d„ Z	d
dej                  dedej                  fd„Zy	)ÚFlaxResNetEmbeddingszO
    ResNet Embeddings (stem) composed of a single aggressive convolution.
    Úconfigr/   c                 óÖ   — t        | j                  j                  dd| j                  j                  | j                  ¬«      | _        t        t        j                  ddd¬«      | _        y )Né   r1   )r+   r,   r.   r/   )r   r   )r1   r1   )©r   r   r]   )Úwindow_shaper7   r8   )	r)   rZ   Úembedding_sizeÚ
hidden_actr/   Úembedderr   r&   Úmax_poolrG   s    r   rH   zFlaxResNetEmbeddings.setup„   sN   € Ü+Ø�K‰K×&Ñ&ØØØ—{‘{×-Ñ-Ø—*‘*ô
ˆŒô  ¤§¡¸&È&ÐZjÔkˆ�r!   Úpixel_valuesrI   rJ   c                 ó´   — |j                   d   }|| j                  j                  k7  rt        d«      ‚| j	                  ||¬«      }| j                  |«      }|S )NéÿÿÿÿzeMake sure that the channel dimension of the pixel values match with the one set in the configuration.©rI   )ÚshaperZ   Únum_channelsÚ
ValueErrorra   rb   )r   rc   rI   rh   Ú	embeddings        r   r    zFlaxResNetEmbeddings.__call__�   s\   € Ø#×)Ñ)¨"Ñ-ˆØ˜4Ÿ;™;×3Ñ3Ò3ÜØwóð ð —M‘M ,¸m�MÓLˆ	Ø—M‘M )Ó,ˆ	ØÐr!   NrP   )r"   r#   r$   r%   r   rR   rT   rU   r/   rH   rV   rW   r    r   r!   r   rY   rY   |   sL   … ñð ÓØ—{‘{€Eˆ3�9‰9Ó"ò	lñ S§[¡[ð Àð ÐQT×Q\ÑQ\ô r!   rY   c                   ó¤   — e Zd ZU dZeed<   dZeed<   ej                  Z	ej                  ed<   d„ Z
ddej                  ded	ej                  fd
„Zy)ÚFlaxResNetShortCutzž
    ResNet shortcut, used to project the residual features to the correct size. If needed, it is also used to
    downsample the input using `stride=2`.
    r*   r1   r,   r/   c                 ó  — t        j                  | j                  d| j                  dt         j                  j                  ddd¬«      | j                  ¬«      | _        t        j                  dd	| j                  ¬
«      | _	        y )Nr]   Fr2   r3   Útruncated_normal)r5   r6   )r+   r7   r9   r:   r/   r;   r<   r=   )
r&   r@   r*   r,   rA   rB   r/   rC   rD   rE   rG   s    r   rH   zFlaxResNetShortCut.setup¤   se   € ÜŸ7™7Ø×ÑØØ—K‘KØÜŸ™×8Ñ8¸À9Ð[mÐ8ÓnØ—*‘*ô
ˆÔô  Ÿ\™\°3ÀÈTÏZÉZÔXˆÕr!   r   rI   rJ   c                 óN   — | j                  |«      }| j                  ||¬«      }|S rL   )rC   rE   rN   s       r   r    zFlaxResNetShortCut.__call__¯   s-   € Ø×'Ñ'¨Ó*ˆØ×)Ñ)¨,ÈMÐ)ÓZˆØÐr!   NrP   )r"   r#   r$   r%   rQ   rR   r,   rT   rU   r/   rH   rV   rW   r    r   r!   r   rl   rl   š   sR   … ñð
 ÓØ€FˆCƒOØ—{‘{€Eˆ3�9‰9Ó"ò	Yñ˜#Ÿ+™+ð °dð ÀcÇkÁkô r!   rl   c                   ó    — e Zd ZU eed<   dZeed<   ej                  Zej                  ed<   d„ Z	ddej                  dedej                  fd	„Zy
)ÚFlaxResNetBasicLayerCollectionr*   r   r,   r/   c                 óª   — t        | j                  | j                  | j                  ¬«      t        | j                  d | j                  ¬«      g| _        y )N©r,   r/   )r.   r/   )r)   r*   r,   r/   ÚlayerrG   s    r   rH   z$FlaxResNetBasicLayerCollection.setupº   s;   € ä × 1Ñ 1¸$¿+¹+ÈTÏZÉZÔXÜ × 1Ñ 1¸dÈ$Ï*É*ÔUð
ˆ�
r!   rO   rI   rJ   c                 ó<   — | j                   D ]  } |||¬«      }Œ |S ©Nrf   ©rt   ©r   rO   rI   rt   s       r   r    z'FlaxResNetBasicLayerCollection.__call__À   ó)   € Ø—Z‘Zò 	LˆEÙ  ¸]ÔK‰Lð	LàÐr!   NrP   )r"   r#   r$   rQ   rR   r,   rT   rU   r/   rH   rV   rW   r    r   r!   r   rq   rq   µ   sM   … ØÓØ€FˆCƒOØ—{‘{€Eˆ3�9‰9Ó"ò
ñ S§[¡[ð Àð ÐQT×Q\ÑQ\ô r!   rq   c                   ó’   — e Zd ZU dZeed<   eed<   dZeed<   dZee	   ed<   e
j                  Ze
j                  ed<   d	„ Zdd
efd„Zy)ÚFlaxResNetBasicLayerzO
    A classic ResNet's residual layer composed by two `3x3` convolutions.
    Úin_channelsr*   r   r,   r-   r.   r/   c                 óT  — | j                   | j                  k7  xs | j                  dk7  }|r,t        | j                  | j                  | j                  ¬«      nd | _        t        | j                  | j                  | j                  ¬«      | _        t        | j                     | _
        y )Nr   rs   )r*   r,   r/   )r|   r*   r,   rl   r/   Úshortcutrq   rt   r   r.   rF   ©r   Úshould_apply_shortcuts     r   rH   zFlaxResNetBasicLayer.setupÑ   s‹   € Ø $× 0Ñ 0°D×4EÑ4EÑ EÒ YÈÏÉÐXYÑIYÐñ %ô ˜t×0Ñ0¸¿¹ÈDÏJÉJÕWàð 	Œô
 4Ø×*Ñ*Ø—;‘;Ø—*‘*ô
ˆŒ
ô
  & d§o¡oÑ6ˆÕr!   rI   c                 óš   — |}| j                  ||¬«      }| j                  �| j                  ||¬«      }||z  }| j                  |«      }|S rv   )rt   r~   rF   ©r   rO   rI   Úresiduals       r   r    zFlaxResNetBasicLayer.__call__ß   sU   € ØˆØ—z‘z ,¸m�zÓLˆà�=‰=Ð$Ø—}‘} X¸]�}ÓKˆHØ˜Ñ ˆà×+Ñ+¨LÓ9ˆØÐr!   NrP   )r"   r#   r$   r%   rQ   rR   r,   r.   r   rS   rT   rU   r/   rH   rW   r    r   r!   r   r{   r{   Æ   sO   … ñð ÓØÓØ€FˆCƒOØ &€J�˜‘Ó&Ø—{‘{€Eˆ3�9‰9Ó"ò7ñ	°Dô 	r!   r{   c                   óÂ   — e Zd ZU eed<   dZeed<   dZee   ed<   dZ	eed<   e
j                  Ze
j                  ed<   d	„ Zdd
e
j                  dede
j                  fd„Zy)Ú#FlaxResNetBottleNeckLayerCollectionr*   r   r,   r-   r.   é   Ú	reductionr/   c           	      óþ   — | j                   | j                  z  }t        |d| j                  d¬«      t        || j                  | j                  d¬«      t        | j                   dd | j                  d¬«      g| _        y )Nr   Ú0)r+   r/   ÚnameÚ1)r,   r/   rŠ   Ú2)r+   r.   r/   rŠ   )r*   r‡   r)   r/   r,   rt   )r   Úreduces_channelss     r   rH   z)FlaxResNetBottleNeckLayerCollection.setupò   sl   € Ø×,Ñ,°·±Ñ>Ðô  Ð 0¸aÀtÇzÁzÐX[Ô\ÜÐ 0¸¿¹ÈDÏJÉJÐ]`ÔaÜ × 1Ñ 1¸qÈTÐY]×YcÑYcÐjmÔnð
ˆ�
r!   rO   rI   rJ   c                 ó<   — | j                   D ]  } |||¬«      }Œ |S rv   rw   rx   s       r   r    z,FlaxResNetBottleNeckLayerCollection.__call__û   ry   r!   NrP   )r"   r#   r$   rQ   rR   r,   r.   r   rS   r‡   rT   rU   r/   rH   rV   rW   r    r   r!   r   r…   r…   ë   se   … ØÓØ€FˆCƒOØ &€J�˜‘Ó&Ø€IˆsÓØ—{‘{€Eˆ3�9‰9Ó"ò
ñ S§[¡[ð Àð ÐQT×Q\ÑQ\ô r!   r…   c                   óÐ   — e Zd ZU dZeed<   eed<   dZeed<   dZee	   ed<   dZ
eed	<   ej                  Zej                  ed
<   d„ Zddej                  dedej                  fd„Zy)ÚFlaxResNetBottleNeckLayera$  
    A classic ResNet's bottleneck layer composed by three `3x3` convolutions. The first `1x1` convolution reduces the
    input by a factor of `reduction` in order to make the second `3x3` convolution faster. The last `1x1` convolution
    remaps the reduced features to `out_channels`.
    r|   r*   r   r,   r-   r.   r†   r‡   r/   c                 ó€  — | j                   | j                  k7  xs | j                  dk7  }|r,t        | j                  | j                  | j                  ¬«      nd | _        t        | j                  | j                  | j                  | j                  | j                  ¬«      | _	        t        | j                     | _        y )Nr   rs   )r,   r.   r‡   r/   )r|   r*   r,   rl   r/   r~   r…   r.   r‡   rt   r   rF   r   s     r   rH   zFlaxResNetBottleNeckLayer.setup  s™   € Ø $× 0Ñ 0°D×4EÑ4EÑ EÒ YÈÏÉÐXYÑIYÐñ %ô ˜t×0Ñ0¸¿¹ÈDÏJÉJÕWàð 	Œô 9Ø×ÑØ—;‘;Ø—‘Ø—n‘nØ—*‘*ô
ˆŒ
ô  & d§o¡oÑ6ˆÕr!   rO   rI   rJ   c                 ó˜   — |}| j                   �| j                  ||¬«      }| j                  ||«      }||z  }| j                  |«      }|S rv   )r~   rt   rF   r‚   s       r   r    z"FlaxResNetBottleNeckLayer.__call__!  sS   € Øˆà�=‰=Ð$Ø—}‘} X¸]�}ÓKˆHØ—z‘z ,°Ó>ˆØ˜Ñ ˆØ×+Ñ+¨LÓ9ˆØÐr!   NrP   )r"   r#   r$   r%   rQ   rR   r,   r.   r   rS   r‡   rT   rU   r/   rH   rV   rW   r    r   r!   r   r�   r�     sr   … ñð ÓØÓØ€FˆCƒOØ &€J�˜‘Ó&Ø€IˆsÓØ—{‘{€Eˆ3�9‰9Ó"ò7ñ$ S§[¡[ð Àð ÐQT×Q\ÑQ\ô r!   r�   c                   óÆ   — e Zd ZU dZ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	„ Zdd
e	j                  dede	j                  fd„Zy)ÚFlaxResNetStageLayersCollectionú4
    A ResNet stage composed by stacked layers.
    rZ   r|   r*   r1   r,   Údepthr/   c                 óà  — | j                   j                  dk(  rt        nt        } || j                  | j
                  | j                  | j                   j                  | j                  d¬«      g}t        | j                  dz
  «      D ]\  }|j                   || j
                  | j
                  | j                   j                  | j                  t        |dz   «      ¬«      «       Œ^ || _        y )NÚ
bottleneckr‰   )r,   r.   r/   rŠ   r   )r.   r/   rŠ   )rZ   Ú
layer_typer�   r{   r|   r*   r,   r`   r/   Úranger–   ÚappendrS   Úlayers)r   rt   rœ   Úis       r   rH   z%FlaxResNetStageLayersCollection.setup8  sÉ   € Ø-1¯[©[×-CÑ-CÀ|Ò-SÕ)ÔYmˆñ Ø× Ñ Ø×!Ñ!Ø—{‘{ØŸ;™;×1Ñ1Ø—j‘jØôð

ˆô �t—z‘z A‘~Ó&ò 		ˆAØ�M‰MÙØ×%Ñ%Ø×%Ñ%Ø#Ÿ{™{×5Ñ5ØŸ*™*Ü˜Q ™U›ôõð		ð ˆ�r!   r   rI   rJ   c                 ó@   — |}| j                   D ]  } |||¬«      }Œ |S rv   ©rœ   )r   r   rI   rO   rt   s        r   r    z(FlaxResNetStageLayersCollection.__call__T  s.   € ØˆØ—[‘[ò 	LˆEÙ  ¸]ÔK‰Lð	LàÐr!   NrP   ©r"   r#   r$   r%   r   rR   rQ   r,   r–   rT   rU   r/   rH   rV   rW   r    r   r!   r   r”   r”   ,  sf   … ñð ÓØÓØÓØ€FˆCƒOØ€Eˆ3ƒNØ—{‘{€Eˆ3�9‰9Ó"òñ8˜#Ÿ+™+ð °dð ÀcÇkÁkô r!   r”   c                   óÆ   — e Zd ZU dZ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	„ Zdd
e	j                  dede	j                  fd„Zy)ÚFlaxResNetStager•   rZ   r|   r*   r1   r,   r–   r/   c                 ó¨   — t        | j                  | j                  | j                  | j                  | j
                  | j                  ¬«      | _        y )N)r|   r*   r,   r–   r/   )r”   rZ   r|   r*   r,   r–   r/   rœ   rG   s    r   rH   zFlaxResNetStage.setupg  s<   € Ü5Ø�K‰KØ×(Ñ(Ø×*Ñ*Ø—;‘;Ø—*‘*Ø—*‘*ô
ˆ�r!   r   rI   rJ   c                 ó(   — | j                  ||¬«      S rv   rŸ   )r   r   rI   s      r   r    zFlaxResNetStage.__call__q  s   € Ø�{‰{˜1¨Mˆ{Ó:Ð:r!   NrP   r    r   r!   r   r¢   r¢   [  sf   … ñð ÓØÓØÓØ€FˆCƒOØ€Eˆ3ƒNØ—{‘{€Eˆ3�9‰9Ó"ò
ñ;˜#Ÿ+™+ð ;°dð ;ÀcÇkÁkô ;r!   r¢   c            	       ó†   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 d
dej                  de
de
defd„Zy	)ÚFlaxResNetStageCollectionrZ   r/   c                 óv  — t        | j                  j                  | j                  j                  dd  «      }t        | j                  | j                  j                  | j                  j                  d   | j                  j
                  rdnd| j                  j                  d   | j                  d¬«      g}t        t        || j                  j                  dd  «      «      D ]K  \  }\  \  }}}|j                  t        | j                  |||| j                  t        |dz   «      ¬«      «       ŒM || _        y )Nr   r   r1   r‰   )r,   r–   r/   rŠ   )r–   r/   rŠ   )ÚziprZ   Úhidden_sizesr¢   r_   Údownsample_in_first_stageÚdepthsr/   Ú	enumerater›   rS   Ústages)r   Úin_out_channelsr­   r�   r|   r*   r–   s          r   rH   zFlaxResNetStageCollection.setupy  s  € Ü˜dŸk™k×6Ñ6¸¿¹×8PÑ8PÐQRÐQSÐ8TÓUˆäØ—‘Ø—‘×*Ñ*Ø—‘×(Ñ(¨Ñ+Ø ŸK™K×AÒA‘qÀqØ—k‘k×(Ñ(¨Ñ+Ø—j‘jØôð

ˆô 8AÄÀ_ÐVZ×VaÑVa×VhÑVhÐijÐikÐVlÓAmÓ7nò 	Ñ3ˆAÑ3Ñ+�˜l¨UØ�M‰MÜ §¡¨[¸,ÈeÐ[_×[eÑ[eÔloÐpqÐtuÑpuÓlvÔwõð	ð
 ˆ�r!   rO   Úoutput_hidden_statesrI   rJ   c                 ó€   — |rdnd }| j                   D ]&  }|r||j                  dddd«      fz   } |||¬«      }Œ( ||fS )Nr   r   r   r   r1   rf   )r­   Ú	transpose)r   rO   r¯   rI   Úhidden_statesÚstage_modules         r   r    z"FlaxResNetStageCollection.__call__Ž  s]   € ñ 3™¸ˆà ŸK™Kò 	SˆLÙ#Ø -°×1GÑ1GÈÈ1ÈaÐQRÓ1SÐ0UÑ U�á'¨ÀMÔR‰Lð		Sð ˜]Ð*Ð*r!   N)FT©r"   r#   r$   r   rR   rT   rU   r/   rH   rV   rW   r   r    r   r!   r   r¦   r¦   u  sV   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òð0 &+Ø"ñ	+à—k‘kð+ð #ð+ð ð	+ð
 
,ô+r!   r¦   c                   óŒ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 ddej                  de
de
de
def
d	„Zy
)ÚFlaxResNetEncoderrZ   r/   c                 óP   — t        | j                  | j                  ¬«      | _        y )N©r/   )r¦   rZ   r/   r­   rG   s    r   rH   zFlaxResNetEncoder.setup£  s   € Ü/°·±À4Ç:Á:ÔNˆ�r!   rO   r¯   Úreturn_dictrI   rJ   c                 óª   — | j                  |||¬«      \  }}|r||j                  dddd«      fz   }|st        d„ ||fD «       «      S t        ||¬«      S )N)r¯   rI   r   r   r   r1   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wr   r   )Ú.0Úvs     r   ú	<genexpr>z-FlaxResNetEncoder.__call__.<locals>.<genexpr>µ  s   è ø€ ÒS˜qÀQÁ]œÑSùs   ‚Š)Úlast_hidden_stater²   )r­   r±   Útupler   )r   rO   r¯   r¹   rI   r²   s         r   r    zFlaxResNetEncoder.__call__¦  st   € ð '+§k¡kØÐ/CÐS`ð '2ó '
Ñ#ˆ�mñ  Ø)¨\×-CÑ-CÀAÀqÈ!ÈQÓ-OÐ,QÑQˆMáÜÑS \°=Ð$AÔSÓSÐSä1Ø*Ø'ô
ð 	
r!   N)FTTr´   r   r!   r   r¶   r¶   Ÿ  sd   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òOð &+Ø Ø"ñ
à—k‘kð
ð #ð
ð ð	
ð
 ð
ð 
,ô
r!   r¶   c                   ó  ‡ — 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«      	 	 	 	 ddededee   dee   fd„«       Zˆ xZS )ÚFlaxResNetPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úresnetrc   NÚmodule_class)r   éà   rÅ   r   r   TrZ   Úseedr/   Ú_do_initc                 ó¦   •—  | j                   d||dœ|¤Ž}|€$d|j                  |j                  |j                  f}t        ‰| �  ||||||¬«       y )N©rZ   r/   r   )Úinput_shaperÆ   r/   rÇ   r   )rÄ   Ú
image_sizerh   ÚsuperÚ__init__)	r   rZ   rÊ   rÆ   r/   rÇ   r   ÚmoduleÚ	__class__s	           €r   rÍ   z"FlaxResNetPreTrainedModel.__init__È  sc   ø€ ð #�×"Ñ"ÐH¨&¸ÑHÀÑHˆØÐØ˜f×/Ñ/°×1BÑ1BÀF×DWÑDWÐXˆKÜ‰Ñ˜ °[ÀtÐSXÐckÐÕlr!   ÚrngrÊ   ÚparamsrJ   c                 óX  — t        j                  || j                  ¬«      }d|i}| j                  j	                  ||d¬«      }|�dt        t        |«      «      }t        t        |«      «      }| j                  D ]
  }||   ||<   Œ t        «       | _        t        t        |«      «      S |S )Nr¸   rÑ   F)r¹   )rT   Úzerosr/   rÎ   Úinitr	   r   Ú_missing_keysÚsetr   r
   )r   rÐ   rÊ   rÑ   rc   ÚrngsÚrandom_paramsÚmissing_keys           r   Úinit_weightsz&FlaxResNetPreTrainedModel.init_weightsÖ  s¤   € ä—y‘y °D·J±JÔ?ˆà˜#ˆˆàŸ™×(Ñ(¨¨|ÈÐ(ÓOˆàÐÜ(¬°-Ó)@ÓAˆMÜ!¤(¨6Ó"2Ó3ˆFØ#×1Ñ1ò A�Ø&3°KÑ&@��{Ò#ðAä!$£ˆDÔÜœ.¨Ó0Ó1Ð1à Ð r!   Útrainr¯   r¹   c           	      ó�  — |�|n| j                   j                  }|�|n| j                   j                  }t        j                  |d«      }i }| j
                  j                  |�|d   n| j                  d   |�|d   n| j                  d   dœt        j                  |t        j                  ¬«      | ||||rdg¬«      S d¬«      S )N)r   r1   r   r   rÑ   Úbatch_stats)rÑ   rÝ   r¸   F)r×   Úmutable)
rZ   r¯   r¹   rT   r±   rÎ   ÚapplyrÑ   ÚarrayrU   )r   rc   rÑ   rÛ   r¯   r¹   r×   s          r   r    z"FlaxResNetPreTrainedModel.__call__è  sà   € ð %9Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆä—}‘} \°<Ó@ˆð ˆà�{‰{× Ñ à.4Ð.@˜& Ò*ÀdÇkÁkÐRZÑF[Ø8>Ð8J˜v mÒ4ÐPT×P[ÑP[Ð\iÑPjñô �I‰I�l¬#¯+©+Ô6ØˆIØ ØØÙ',�]�Oð !ó 
ð 	
ð 38ð !ó 
ð 	
r!   r   )NFNN)r"   r#   r$   r%   r   Úconfig_classÚbase_model_prefixÚmain_input_namerÄ   r&   ÚModulerR   rT   rU   rQ   r/   rW   rÍ   ÚjaxÚrandomÚPRNGKeyr   r   rÚ   r   ÚRESNET_INPUTS_DOCSTRINGÚdictr   r    Ú__classcell__)rÏ   s   @r   rÂ   rÂ   ½  sô   ø… ñð
  €LØ ÐØ$€OØ"€L�"—)‘)Ó"ð
 %ØØŸ;™;Øñmàðmð ð	mð
 �y‰yðmð õmñ! §
¡
× 2Ñ 2ð !Àð !ÐPZð !Ðfpó !ñ$ +Ð+BÓCð ØØ/3Ø&*ñ
ð ð
ð ð	
ð
 ' t™nð
ð ˜d‘^ò
ó Dô
r!   rÂ   c            	       ót   — 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	)ÚFlaxResNetModulerZ   r/   c                 óÜ   — t        | j                  | j                  ¬«      | _        t	        | j                  | j                  ¬«      | _        t        t        j                  d¬«      | _	        y )Nr¸   )©r   r   rî   )r8   )
rY   rZ   r/   ra   r¶   Úencoderr   r&   Úavg_poolÚpoolerrG   s    r   rH   zFlaxResNetModule.setup  sF   € Ü,¨T¯[©[ÀÇ
Á
ÔKˆŒÜ(¨¯©¸D¿J¹JÔGˆŒô Ü�K‰KØ$ô
ˆ�r!   rI   r¯   r¹   rJ   c                 óð  — |�|n| j                   j                  }|�|n| j                   j                  }| j                  ||¬«      }| j	                  ||||¬«      }|d   }| j                  ||j                  d   |j                  d   f|j                  d   |j                  d   f¬«      j                  dddd«      }|j                  dddd«      }|s
||f|dd  z   S t        |||j                  ¬«      S )	Nrf   )r¯   r¹   rI   r   r   r1   )r^   r7   r   )r¿   Úpooler_outputr²   )
rZ   r¯   Úuse_return_dictra   rï   rñ   rg   r±   r   r²   )	r   rc   rI   r¯   r¹   Úembedding_outputÚencoder_outputsr¿   Úpooled_outputs	            r   r    zFlaxResNetModule.__call__  s-  € ð %9Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàŸ=™=¨À]˜=ÓSÐàŸ,™,ØØ!5Ø#Ø'ð	 'ó 
ˆð ,¨AÑ.ÐàŸ™ØØ+×1Ñ1°!Ñ4Ð6G×6MÑ6MÈaÑ6PÐQØ&×,Ñ,¨QÑ/Ð1B×1HÑ1HÈÑ1KÐLð $ó 
÷ ‰)�A�q˜!˜QÓ
ð	 	ð .×7Ñ7¸¸1¸aÀÓCÐáØ% }Ð5¸ÈÈÐ8KÑKÐKä;Ø/Ø'Ø)×7Ñ7ô
ð 	
r!   N)TFT)r"   r#   r$   r   rR   rT   rU   r/   rH   rW   r   r    r   r!   r   rì   rì   	  sW   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð #Ø%*Ø ñ&
ð ð&
ð #ð	&
ð
 ð&
ð 
6ô&
r!   rì   zOThe bare ResNet model outputting raw features without any specific head on top.c                   ó   — e Zd ZeZy)ÚFlaxResNetModelN)r"   r#   r$   rì   rÄ   r   r!   r   rù   rù   @  s	   „ ð
 $�Lr!   rù   an  
    Returns:

    Examples:

    ```python
    >>> from transformers import AutoImageProcessor, FlaxResNetModel
    >>> 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("microsoft/resnet-50")
    >>> model = FlaxResNetModel.from_pretrained("microsoft/resnet-50")
    >>> inputs = image_processor(images=image, return_tensors="np")
    >>> outputs = model(**inputs)
    >>> last_hidden_states = outputs.last_hidden_state
    ```
)Úoutput_typerá   c                   óŒ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zdej                  dej                  fd„Z
y)ÚFlaxResNetClassifierCollectionrZ   r/   c                 óz   — t        j                  | j                  j                  | j                  d¬«      | _        y )Nr‹   )r/   rŠ   )r&   ÚDenserZ   Ú
num_labelsr/   Ú
classifierrG   s    r   rH   z$FlaxResNetClassifierCollection.setupf  s%   € ÜŸ(™( 4§;¡;×#9Ñ#9ÀÇÁÐRUÔVˆ�r!   r   rJ   c                 ó$   — | j                  |«      S r   )r   )r   r   s     r   r    z'FlaxResNetClassifierCollection.__call__i  s   € Ø�‰˜qÓ!Ð!r!   N)r"   r#   r$   r   rR   rT   rU   r/   rH   rV   r    r   r!   r   rü   rü   b  s;   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òWð"˜#Ÿ+™+ð "¨#¯+©+ô "r!   rü   c                   ój   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 dde	fd„Z
y)Ú&FlaxResNetForImageClassificationModulerZ   r/   c                 óî   — t        | j                  | j                  ¬«      | _        | j                  j                  dkD  r't        | j                  | j                  ¬«      | _        y t        «       | _        y )NrÉ   r   r¸   )rì   rZ   r/   rÃ   rÿ   rü   r   r   rG   s    r   rH   z,FlaxResNetForImageClassificationModule.setupq  sL   € Ü&¨d¯k©kÀÇÁÔLˆŒà�;‰;×!Ñ! AÒ%Ü<¸T¿[¹[ÐPT×PZÑPZÔ[ˆD�Oä&›jˆD�Or!   NrI   c                 ó  — |�|n| j                   j                  }| j                  ||||¬«      }|r|j                  n|d   }| j	                  |d d …d d …ddf   «      }|s|f|dd  z   }|S t        ||j                  ¬«      S )N)rI   r¯   r¹   r   r   r1   )Úlogitsr²   )rZ   rô   rÃ   ró   r   r   r²   )	r   rc   rI   r¯   r¹   Úoutputsr÷   r  Úoutputs	            r   r    z/FlaxResNetForImageClassificationModule.__call__y  s—   € ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—+‘+ØØ'Ø!5Ø#ð	 ó 
ˆñ 2=˜×-Ò-À'È!Á*ˆà—‘ ªq²!°Q¸¨zÑ!:Ó;ˆáØ�Y ¨¨ Ñ,ˆFØˆMä7¸vÐU\×UjÑUjÔkÐkr!   )NTNN)r"   r#   r$   r   rR   rT   rU   r/   rH   rW   r    r   r!   r   r  r  m  s>   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò)ð Ø"Ø!Øñlð ôlr!   r  z†
    ResNet Model with an image classification head on top (a linear layer on top of the pooled features), e.g. for
    ImageNet.
    c                   ó   — e Zd ZeZy)Ú FlaxResNetForImageClassificationN)r"   r#   r$   r  rÄ   r   r!   r   r
  r
  ”  s	   „ ð :�Lr!   r
  a]  
    Returns:

    Example:

    ```python
    >>> from transformers import AutoImageProcessor, FlaxResNetForImageClassification
    >>> 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("microsoft/resnet-50")
    >>> model = FlaxResNetForImageClassification.from_pretrained("microsoft/resnet-50")

    >>> 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Â   )9Ú	functoolsr   Útypingr   r   Ú
flax.linenÚlinenr&   rå   Ú	jax.numpyÚnumpyrT   Úflax.core.frozen_dictr   r   r   Úflax.traverse_utilr	   r
   Úmodeling_flax_outputsr   r   r   Úmodeling_flax_utilsr   r   r   r   Úutilsr   r   Úconfiguration_resnetr   ÚRESNET_START_DOCSTRINGrè   rä   r   r)   rY   rl   rq   r{   r…   r�   r”   r¢   r¦   r¶   rÂ   rì   rù   ÚFLAX_VISION_MODEL_DOCSTRINGrü   r  r
  ÚFLAX_VISION_CLASSIF_DOCSTRINGÚ__all__r   r!   r   ú<module>r     sü  ðõ  ß "å Û 
Ý ß >Ñ >ß ;÷ñ ÷
ó ÷ QÝ .ð!Ð ðH
Ð ôˆr�y‰yô ô˜"Ÿ)™)ô ô6˜2Ÿ9™9ô ô<˜Ÿ™ô ô6 R§Y¡Yô ô""˜2Ÿ9™9ô "ôJ¨"¯)©)ô ô,( §	¡	ô (ôV, b§i¡iô ,ô^;�b—i‘iô ;ô4'+ §	¡	ô '+ôT
˜Ÿ	™	ô 
ô<I
Ð 3ô I
ôX4
�r—y‘yô 4
ñn ØUØóô$Ð/ó $ó	ð$ðÐ ñ( ˜Ð*EÔ FÙ  ØÐ!MÐ\hõô
" R§Y¡Yô "ô$l¨R¯Y©Yô $lñN ðð óô:Ð'@ó :óð:ð!Ð ñ6 Ð9Ð;XÔ YÙ  Ø$Ð2ZÐiuõò
 _�r!   