Ë
    S^(h<Æ  ã                   óF  — d dl mZ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mZmZ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 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$ ddl%m&Z&m'Z'm(Z(  e$jR                  e*«      Z+dZ,dZ-dZ.dZ/ej`                  jb                   G d„ de"«      «       Z2ej`                  jb                   G d„ de"«      «       Z3 G d„ dejh                  «      Z5 G d„ dejh                  «      Z6 G d„ dejh                  «      Z7 G d„ dejh                  «      Z8 G d„ dejh                  «      Z9 G d „ d!ejh                  «      Z: G d"„ d#ejh                  «      Z; G d$„ d%ejh                  «      Z< G d&„ d'ejh                  «      Z= G d(„ d)e«      Z> G d*„ d+e«      Z? G d,„ d-e«      Z@ G d.„ d/ejh                  «      ZA G d0„ d1e>«      ZBd2ZC e eBe-eCz   «        eeBee'¬3«        G d4„ d5ejh                  «      ZD G d6„ d7e>«      ZEd8ZF e eEe-eFz   «        eeEe2e'¬3«        G d9„ d:ejh                  «      ZG G d;„ d<e?«      ZHd=ZI e eHe.eIz   «        eeHee(¬3«        G d>„ d?ejh                  «      ZJ e#e,«       G d@„ dAe@«      «       ZKdBZL e eKe/eLz   «        eeKe3e&¬3«       g dC¢ZMy)Dé    )ÚAnyÚOptionalÚTupleÚUnionN)Ú
FrozenDictÚfreezeÚunfreeze)Úcombine_masksÚmake_causal_mask)Údot_product_attention_weights)Úflatten_dictÚunflatten_dict)Úlaxé   )ÚFlaxBaseModelOutputÚFlaxBaseModelOutputWithPooling)ÚACT2FNÚFlaxPreTrainedModelÚ append_replace_return_docstringsÚoverwrite_call_docstring)ÚModelOutputÚadd_start_docstringsÚloggingé   )Ú
CLIPConfigÚCLIPTextConfigÚCLIPVisionConfigaü  

    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 ([`CLIPConfig`]): 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_ids (`numpy.ndarray` of shape `(batch_size, sequence_length)`):
            Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
            it.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            [What are input IDs?](../glossary#input-ids)
        attention_mask (`numpy.ndarray` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to avoid performing 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)
        position_ids (`numpy.ndarray` of shape `(batch_size, sequence_length)`, *optional*):
            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.

            [What are position IDs?](../glossary#position-ids)
        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.
aA  
    Args:
        pixel_values (`numpy.ndarray` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Padding will be ignored by default should you provide it. Pixel values can be obtained using
            [`AutoImageProcessor`]. See [`CLIPImageProcessor.__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.
a§  
    Args:
        input_ids (`numpy.ndarray` of shape `(batch_size, sequence_length)`):
            Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
            it.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            [What are input IDs?](../glossary#input-ids)
        attention_mask (`numpy.ndarray` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to avoid performing 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)
        position_ids (`numpy.ndarray` of shape `(batch_size, sequence_length)`, *optional*):
            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.

            [What are position IDs?](../glossary#position-ids)
        pixel_values (`numpy.ndarray` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Padding will be ignored by default should you provide it. Pixel values can be obtained using
            [`AutoImageProcessor`]. See [`CLIPImageProcessor.__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                   óº   — 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                  df      ed<   dZe
eej                  df      ed<   y)ÚFlaxCLIPTextModelOutputaJ  
    Base class for text model's outputs that also contains a pooling of the last hidden states.

    Args:
        text_embeds (`jnp.ndarray` of shape `(batch_size, output_dim`):
            The text embeddings obtained by applying the projection layer to the pooled output of
            [`FlaxCLIPTextModel`].
        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.
        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Útext_embedsÚlast_hidden_state.Úhidden_statesÚ
attentions)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r    ÚjnpÚndarrayÚ__annotations__r!   r"   r   r   r#   © ó    úi/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/clip/modeling_flax_clip.pyr   r   Ÿ   s`   … ñð,  $€K�—‘Ó#Ø%)Ð�s—{‘{Ó)Ø7;€M�8˜E #§+¡+¨sÐ"2Ñ3Ñ4Ó;Ø48€J�˜˜sŸ{™{¨CÐ/Ñ0Ñ1Ô8r,   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j                  ed<   dZeed<   dZeed<   d	ee   fd
„Zy)ÚFlaxCLIPOutputah  
    Args:
        logits_per_image:(`jnp.ndarray` of shape `(image_batch_size, text_batch_size)`):
            The scaled dot product scores between `image_embeds` and `text_embeds`. This represents the image-text
            similarity scores.
        logits_per_text:(`jnp.ndarray` of shape `(text_batch_size, image_batch_size)`):
            The scaled dot product scores between `text_embeds` and `image_embeds`. This represents the text-image
            similarity scores.
        text_embeds(`jnp.ndarray` of shape `(batch_size, output_dim`):
            The text embeddings obtained by applying the projection layer to the pooled output of
            [`FlaxCLIPTextModel`].
        image_embeds(`jnp.ndarray` of shape `(batch_size, output_dim`):
            The image embeddings obtained by applying the projection layer to the pooled output of
            [`FlaxCLIPVisionModel`].
        text_model_output(`FlaxBaseModelOutputWithPooling`):
            The output of the [`FlaxCLIPTextModel`].
        vision_model_output(`FlaxBaseModelOutputWithPooling`):
            The output of the [`FlaxCLIPVisionModel`].
    NÚlogits_per_imageÚlogits_per_textr    Úimage_embedsÚtext_model_outputÚvision_model_outputÚreturnc                 óH   ‡ — t        ˆ fd„‰ j                  «       D «       «      S )Nc              3   ód   •K  — | ]'  }|d vr‰|   nt        ‰|«      j                  «       –— Œ) y­w))r3   r4   N)ÚgetattrÚto_tuple)Ú.0ÚkÚselfs     €r-   ú	<genexpr>z*FlaxCLIPOutput.to_tuple.<locals>.<genexpr>Û   s=   øè ø€ ò 
àð Ð LÑLˆD�ŠGÔRYÐZ^Ð`aÓRb×RkÑRkÓRmÓmñ
ùs   ƒ-0)ÚtupleÚkeys©r<   s   `r-   r9   zFlaxCLIPOutput.to_tupleÚ   s#   ø€ Üó 
à—Y‘Y“[ô
ó 
ð 	
r,   )r$   r%   r&   r'   r0   r(   r)   r*   r1   r    r2   r3   r   r4   r   r   r9   r+   r,   r-   r/   r/   ½   sj   … ñð( %)Ð�c—k‘kÓ(Ø#'€O�S—[‘[Ó'Ø#€K�—‘Ó#Ø $€L�#—+‘+Ó$Ø8<ÐÐ5Ó<Ø:>ÐÐ7Ó>ð
˜% ™*ô 
r,   r/   c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)ÚFlaxCLIPVisionEmbeddingsÚconfigÚdtypec           
      óÖ  — | j                   j                  }| j                   j                  }| j                   j                  }| j	                  dt
        j                  j                  j                  d¬«      |f«      | _	        t        j                  |||f||fdd| j                  t
        j                  j                  j                  «       ¬«      | _        ||z  dz  | _        | j                  dz   }t        j                  ||t
        j                  j                  j                  «       ¬	«      | _        t!        j"                  t!        j$                  d
|d¬«      d
¬«      | _        y )NÚclass_embeddingç{®Gáz”?)ÚstddevÚVALIDF)Úkernel_sizeÚstridesÚpaddingÚuse_biasrD   Úkernel_inité   r   ©Úembedding_initr   Úi4©rD   ©Úaxis)rC   Úhidden_sizeÚ
image_sizeÚ
patch_sizeÚparamÚjaxÚnnÚinitializersÚnormalrF   ÚConvrD   Úpatch_embeddingÚnum_patchesÚEmbedÚposition_embeddingr(   Úexpand_dimsÚarangeÚposition_ids)r<   Ú	embed_dimrW   rX   Únum_positionss        r-   ÚsetupzFlaxCLIPVisionEmbeddings.setupå   s  € Ø—K‘K×+Ñ+ˆ	Ø—[‘[×+Ñ+ˆ
Ø—[‘[×+Ñ+ˆ
à#Ÿz™zÐ*;¼S¿V¹V×=PÑ=P×=WÑ=WÐ_cÐ=WÓ=dÐgpÐfrÓsˆÔä!Ÿw™wØØ# ZÐ0Ø Ð,ØØØ—*‘*ÜŸ™×+Ñ+×2Ñ2Ó4ô 
ˆÔð '¨*Ñ4¸Ñ:ˆÔØ×(Ñ(¨1Ñ,ˆÜ"$§(¡(¨=¸)ÔTW×TZÑTZ×TgÑTg×TnÑTnÓTpÔ"qˆÔÜŸO™O¬C¯J©J°q¸-ÈtÔ,TÐ[\Ô]ˆÕr,   c                 ód  — | j                  |«      }|j                  \  }}}}t        j                  ||||z  |f«      }t        j                  | j
                  d¬«      }t        j                  ||ddf«      }t        j                  ||gd¬«      }|| j                  | j                  «      z   }|S )N©r   r   rT   r   )
r_   Úshaper(   Úreshaperc   rF   ÚtileÚconcatenaterb   re   )	r<   Úpixel_valuesÚpatch_embedsÚ
batch_sizeÚheightÚwidthÚchannelsÚclass_embedsÚ
embeddingss	            r-   Ú__call__z!FlaxCLIPVisionEmbeddings.__call__û   s¤   € Ø×+Ñ+¨LÓ9ˆØ.:×.@Ñ.@Ñ+ˆ
�F˜E 8Ü—{‘{ <°*¸fÀu¹nÈhÐ1WÓXˆä—‘ t×';Ñ';À&ÔIˆÜ—x‘x ¨z¸1¸aÐ.@ÓAˆÜ—_‘_ l°LÐ%AÈÔJˆ
Ø $×"9Ñ"9¸$×:KÑ:KÓ"LÑLˆ
ØÐr,   N)
r$   r%   r&   r   r*   r(   Úfloat32rD   rh   rw   r+   r,   r-   rB   rB   á   s%   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò^ó,	r,   rB   c                   óZ   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Zd„ Z	y)ÚFlaxCLIPTextEmbeddingsrC   rD   c                 ó  — | j                   j                  }t        j                  | j                   j                  |t
        j                  j                  j                  «       ¬«      | _        t        j                  | j                   j                  |t
        j                  j                  j                  «       ¬«      | _
        t        j                  t        j                  d| j                   j                  d¬«      d¬«      | _        y )NrP   r   rR   rS   rj   rT   )rC   rV   r[   ra   Ú
vocab_sizerZ   r\   r]   Útoken_embeddingÚmax_position_embeddingsrb   r(   rc   rd   re   )r<   rf   s     r-   rh   zFlaxCLIPTextEmbeddings.setup  s«   € Ø—K‘K×+Ñ+ˆ	ä!Ÿx™x¨¯©×(>Ñ(>À	ÔZ]×Z`ÑZ`×ZmÑZm×ZtÑZtÓZvÔwˆÔÜ"$§(¡(Ø�K‰K×/Ñ/°Ì3Ï6É6×K^ÑK^×KeÑKeÓKgô#
ˆÔô  ŸO™OÜ�J‰J�q˜$Ÿ+™+×=Ñ=ÀTÔJÐQWô
ˆÕr,   c                 ó�   — | j                  |j                  d«      «      }| j                  |j                  d«      «      }||z   }|S )NrR   )r}   Úastyperb   )r<   Ú	input_idsre   Úinput_embedsÚposition_embedsrv   s         r-   rw   zFlaxCLIPTextEmbeddings.__call__  sH   € Ø×+Ñ+¨I×,<Ñ,<¸TÓ,BÓCˆØ×1Ñ1°,×2EÑ2EÀdÓ2KÓLˆà! OÑ3ˆ
ØÐr,   N)
r$   r%   r&   r   r*   r(   rx   rD   rh   rw   r+   r,   r-   rz   rz     s$   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò	
ór,   rz   c                   ó‚   — e Zd ZU eeef   ed<   ej                  Z	ej                  ed<   d„ Z
d„ Zd„ Z	 	 	 d
dedefd	„Zy)ÚFlaxCLIPAttentionrC   rD   c                 ó0  — | j                   j                  | _        | j                   j                  | _        | j                  | j                  z  | _        | j
                  | j                  z  | j                  k7  r&t        d| j                  › d| j                  › d�«      ‚| j
                  dz  | _        | j                   j                  | _	        t        j                  | j                  | j                  t        j                  j                  j                  d«      ¬«      | _        t        j                  | j                  | j                  t        j                  j                  j                  d«      ¬«      | _        t        j                  | j                  | j                  t        j                  j                  j                  d«      ¬«      | _        t        j                  | j                  | j                  t        j                  j                  j                  d«      ¬«      | _        t)        | j                   t*        «      | _        | j,                  r<t/        t1        j2                  d| j                   j4                  fd¬	«      «      | _        y y )
Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: z).g      à¿ç{®Gáz„?©rD   rN   r   rR   rS   )rC   rV   rf   Únum_attention_headsÚ	num_headsÚhead_dimÚ
ValueErrorÚscaleÚattention_dropoutÚdropoutr[   ÚDenserD   rZ   r\   r]   Úk_projÚv_projÚq_projÚout_projÚ
isinstancer   Úcausalr   r(   Úonesr~   Úcausal_maskr@   s    r-   rh   zFlaxCLIPAttention.setup"  s²  € ØŸ™×0Ñ0ˆŒØŸ™×8Ñ8ˆŒØŸ™¨$¯.©.Ñ8ˆŒØ�=‰=˜4Ÿ>™>Ñ)¨T¯^©^Ò;ÜØMÈdÏnÉnÐM]ð ^Ø—N‘NÐ# 2ð'óð ð —]‘] DÑ(ˆŒ
Ø—{‘{×4Ñ4ˆŒä—h‘h˜tŸ~™~°T·Z±ZÌSÏVÉV×M`ÑM`×MgÑMgÐhlÓMmÔnˆŒÜ—h‘h˜tŸ~™~°T·Z±ZÌSÏVÉV×M`ÑM`×MgÑMgÐhlÓMmÔnˆŒÜ—h‘h˜tŸ~™~°T·Z±ZÌSÏVÉV×M`ÑM`×MgÑMgÐhlÓMmÔnˆŒÜŸ™ §¡°t·z±zÌsÏvÉv×ObÑOb×OiÑOiÐjnÓOoÔpˆŒä  §¡¬nÓ=ˆŒØ�;Š;Ü/´·±¸!¸T¿[¹[×=`Ñ=`Ð9aÐimÔ0nÓoˆDÕð r,   c                 óp   — |j                  |j                  d d | j                  | j                  fz   «      S ©NrO   )rl   rk   rŠ   r‹   ©r<   r"   s     r-   Ú_split_headszFlaxCLIPAttention._split_heads7  s5   € Ø×$Ñ$ ]×%8Ñ%8¸¸!Ð%<ÀÇÁÐPT×P]ÑP]Ð?^Ñ%^Ó_Ð_r,   c                 óZ   — |j                  |j                  d d | j                  fz   «      S rš   )rl   rk   rf   r›   s     r-   Ú_merge_headszFlaxCLIPAttention._merge_heads:  s,   € Ø×$Ñ$ ]×%8Ñ%8¸¸!Ð%<ÀÇÁÐ?PÑ%PÓQÐQr,   NÚdeterministicÚoutput_attentionsc           
      ó|  — | j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }d }| j                  r<|j
                  d   |j
                  d   }
}	| j                  d d …d d …|
|	z
  |
…d |
…f   }|�(|�&t        j                  |d¬«      }t        ||d¬«      }n|�|}n|�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"                  || j                  d ¬	«      }t        j(                  d
||«      }| j+                  |«      }| j-                  |«      }|r||f}|S |f}|S )Nr   )éýÿÿÿéþÿÿÿrT   rR   rS   r   g        r�   )ÚbiasÚdropout_rngÚdropout_raterŸ   rD   Ú	precisionz...hqk,...khd->...qhd)r“   r‘   r’   rœ   r–   rk   r˜   r(   rc   r
   r   ÚselectÚfullr€   rD   ÚfinfoÚminr�   Úmake_rngr   Úeinsumrž   r”   )r<   r"   Úattention_maskrŸ   r    ÚqueryÚkeyÚvalueÚcausal_attention_maskÚquery_lengthÚ
key_lengthÚattention_biasr¥   Úattn_weightsÚattn_outputÚoutputss                   r-   rw   zFlaxCLIPAttention.__call__=  s  € ð —‘˜MÓ*ˆØ�k‰k˜-Ó(ˆØ—‘˜MÓ*ˆà×!Ñ! %Ó(ˆØ×Ñ Ó$ˆØ×!Ñ! %Ó(ˆà $ÐØ�;Š;Ø',§{¡{°1¡~°s·y±yÀ±|˜*ˆLØ$(×$4Ñ$4²Qº¸:ÈÑ;TÐWaÐ;aÐcnÐdnÐcnÐ5nÑ$oÐ!àÐ%Ð*?Ð*KÜ Ÿ_™_¨^À(ÔKˆNÜ*¨>Ð;PÐX\Ô]‰NØ"Ð.Ø2‰NØÐ'Ü Ÿ_™_¨^À(Ô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¸,ÈÓNˆØ×'Ñ'¨Ó4ˆØ—m‘m KÓ0ˆá1B�; Ð-ˆØˆð JUÈˆØˆr,   )NTF)r$   r%   r&   r   r   r   r*   r(   rx   rD   rh   rœ   rž   Úboolrw   r+   r,   r-   r…   r…     s[   … Ø�.Ð"2Ð2Ñ3Ó3Ø—{‘{€Eˆ3�9‰9Ó"òpò*`òRð Ø"Ø"'ñ9ð ð	9ð
  ô9r,   r…   c                   ód   — e Zd ZU eeef   ed<   ej                  Z	ej                  ed<   d„ Z
d„ Zy)ÚFlaxCLIPMLPrC   rD   c                 óÐ  — t         | j                  j                     | _        t	        j
                  | j                  j                  | j                  t        j                  j                  j                  d«      ¬«      | _        t	        j
                  | j                  j                  | j                  t        j                  j                  j                  d«      ¬«      | _        y )Nr‡   rˆ   )r   rC   Ú
hidden_actÚactivation_fnr[   r�   Úintermediate_sizerD   rZ   r\   r]   Úfc1rV   Úfc2r@   s    r-   rh   zFlaxCLIPMLP.setup}  s’   € Ü# D§K¡K×$:Ñ$:Ñ;ˆÔÜ—8‘8Ø�K‰K×)Ñ)Ø—*‘*ÜŸ™×+Ñ+×2Ñ2°4Ó8ô
ˆŒô
 —8‘8˜DŸK™K×3Ñ3¸4¿:¹:ÔSV×SYÑSY×SfÑSf×SmÑSmÐnrÓSsÔtˆ�r,   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S ©N)rÀ   r¾   rÁ   r›   s     r-   rw   zFlaxCLIPMLP.__call__†  s4   € ØŸ™ Ó/ˆØ×*Ñ*¨=Ó9ˆØŸ™ Ó/ˆØÐr,   N)r$   r%   r&   r   r   r   r*   r(   rx   rD   rh   rw   r+   r,   r-   r»   r»   y  s0   … Ø�.Ð"2Ð2Ñ3Ó3Ø—{‘{€Eˆ3�9‰9Ó"òuór,   r»   c                   ót   — e Zd ZU eeef   ed<   ej                  Z	ej                  ed<   d„ Z
	 	 ddedefd„Zy)	ÚFlaxCLIPEncoderLayerrC   rD   c                 ó„  — t        | j                  | j                  ¬«      | _        t	        j
                  | j                  j                  | j                  ¬«      | _        t        | j                  | j                  ¬«      | _	        t	        j
                  | j                  j                  | j                  ¬«      | _
        y ©NrS   )ÚepsilonrD   )r…   rC   rD   Ú	self_attnr[   Ú	LayerNormÚlayer_norm_epsÚlayer_norm1r»   ÚmlpÚlayer_norm2r@   s    r-   rh   zFlaxCLIPEncoderLayer.setup‘  sv   € Ü*¨4¯;©;¸d¿j¹jÔIˆŒÜŸ<™<°·±×0JÑ0JÐRV×R\ÑR\Ô]ˆÔÜ˜tŸ{™{°$·*±*Ô=ˆŒÜŸ<™<°·±×0JÑ0JÐRV×R\ÑR\Ô]ˆÕr,   rŸ   r    c                 óÖ   — |}| j                  |«      }| j                  ||||¬«      }|d   }||z   }|}| j                  |«      }| j                  |«      }||z   }|f}|r||dd  z  }|S )N)r"   r®   rŸ   r    r   r   )rÌ   rÉ   rÎ   rÍ   )r<   r"   r®   rŸ   r    ÚresidualÚattn_outputsr¸   s           r-   rw   zFlaxCLIPEncoderLayer.__call__—  s›   € ð !ˆà×(Ñ(¨Ó7ˆØ—~‘~Ø'Ø)Ø'Ø/ð	 &ó 
ˆð % Q™ˆØ  =Ñ0ˆà ˆØ×(Ñ(¨Ó7ˆØŸ™ Ó/ˆØ  =Ñ0ˆà Ð"ˆáØ�| A BÐ'Ñ'ˆGàˆr,   N)TF©r$   r%   r&   r   r   r   r*   r(   rx   rD   rh   r¹   rw   r+   r,   r-   rÅ   rÅ   �  sL   … Ø�.Ð"2Ð2Ñ3Ó3Ø—{‘{€Eˆ3�9‰9Ó"ò^ð #Ø"'ñð ð	ð
  ôr,   rÅ   c            	       ó‚   — e Zd ZU eeef   ed<   ej                  Z	ej                  ed<   d„ Z
	 	 	 	 	 d
dedededefd	„Zy)ÚFlaxCLIPLayerCollectionrC   rD   c           	      óÄ   — t        | j                  j                  «      D �cg c]-  }t        | j                  t	        |«      | j
                  ¬«      ‘Œ/ c}| _        y c c}w )N)ÚnamerD   )ÚrangerC   Únum_hidden_layersrÅ   ÚstrrD   Úlayers)r<   Úis     r-   rh   zFlaxCLIPLayerCollection.setup»  sG   € ô ˜4Ÿ;™;×8Ñ8Ó9ö
àô ! §¡´3°q³6ÀÇÁÖLò
ˆ�ùò 
s   ¢2ANrŸ   r    Úoutput_hidden_statesÚreturn_dictc                 óà   — |rdnd }|rdnd }| 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+   )r:   Úvs     r-   r=   z3FlaxCLIPLayerCollection.__call__.<locals>.<genexpr>ß  s   è ø€ Ò=˜q¨q©}œÑ=ùs   ‚Š)r!   r"   r#   )rÚ   r>   r   )r<   r"   r®   rŸ   r    rÜ   rÝ   Úall_attentionsÚall_hidden_statesÚlayerÚlayer_outputsr¸   s               r-   rw   z FlaxCLIPLayerCollection.__call__Á  sµ   € ñ  1™°dˆÙ"6™B¸DÐà—[‘[ò 
	6ˆEÙ#Ø! mÐ%5Ñ5Ð!á!Ø˜~¸]Ð^oôˆMð *¨!Ñ,ˆMâ Ø =°Ñ#3Ð"5Ñ5‘ð
	6ñ  Ø -Ð!1Ñ1Ðà Ð"ˆáÜÑ= GÔ=Ó=Ð=ä"Ø+Ð;LÐYgô
ð 	
r,   ©NTFFTrÒ   r+   r,   r-   rÔ   rÔ   ·  sh   … Ø�.Ð"2Ð2Ñ3Ó3Ø—{‘{€Eˆ3�9‰9Ó"ò
ð Ø"Ø"'Ø%*Ø ñ"
ð ð	"
ð
  ð"
ð #ð"
ð ô"
r,   rÔ   c            	       ó‚   — e Zd ZU eeef   ed<   ej                  Z	ej                  ed<   d„ Z
	 	 	 	 	 d
dedededefd	„Zy)ÚFlaxCLIPEncoderrC   rD   c                 óP   — t        | j                  | j                  ¬«      | _        y ©NrS   )rÔ   rC   rD   rÚ   r@   s    r-   rh   zFlaxCLIPEncoder.setupê  s   € Ü-¨d¯k©kÀÇÁÔLˆ�r,   NrŸ   r    rÜ   rÝ   c                 ó0   — | j                  ||||||¬«      S )N)r"   r®   rŸ   r    rÜ   rÝ   )rÚ   )r<   Úinputs_embedsr®   rŸ   r    rÜ   rÝ   s          r-   rw   zFlaxCLIPEncoder.__call__í  s,   € ð �{‰{Ø'Ø)Ø'Ø/Ø!5Ø#ð ó 
ð 	
r,   rå   rÒ   r+   r,   r-   rç   rç   æ  si   … Ø�.Ð"2Ð2Ñ3Ó3Ø—{‘{€Eˆ3�9‰9Ó"òMð Ø"Ø"'Ø%*Ø ñ
ð ð	
ð
  ð
ð #ð
ð ô
r,   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	)ÚFlaxCLIPTextTransformerrC   rD   c                 óF  — t        | j                  | j                  ¬«      | _        t	        | j                  | j                  ¬«      | _        t        j                  | j                  j                  | j                  ¬«      | _	        | j                  j                  | _
        y rÇ   )rz   rC   rD   rv   rç   Úencoderr[   rÊ   rË   Úfinal_layer_normÚeos_token_idr@   s    r-   rh   zFlaxCLIPTextTransformer.setup  sf   € Ü0°·±ÀDÇJÁJÔOˆŒÜ& t§{¡{¸$¿*¹*ÔEˆŒÜ "§¡°T·[±[×5OÑ5OÐW[×WaÑWaÔ bˆÔð !ŸK™K×4Ñ4ˆÕr,   rŸ   r    rÜ   rÝ   c                 ó’  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }| j	                  ||¬«      }| j                  ||||||¬«      }	|	d   }
| j                  |
«      }
| j                  dk(  r8|
t        j                  |
j                  d   «      |j                  d¬«      f   }nD|
t        j                  |
j                  d   «      || j                  k(  j                  d¬«      f   }|s
|
|f|	dd  z   S t        |
||	j                  |	j                  ¬«      S )	N)r�   re   )rë   r®   rŸ   r    rÜ   rÝ   r   rO   éÿÿÿÿrT   r   ©r!   Úpooler_outputr"   r#   )rC   r    rÜ   Úuse_return_dictrv   rï   rð   rñ   r(   rd   rk   Úargmaxr   r"   r#   )r<   r�   r®   re   rŸ   r    rÜ   rÝ   r"   Úencoder_outputsr!   Úpooled_outputs               r-   rw   z FlaxCLIPTextTransformer.__call__  sq  € ð 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàŸ™°)È,˜ÓWˆàŸ,™,Ø'Ø)Ø'Ø/Ø!5Ø#ð 'ó 
ˆð ,¨AÑ.ÐØ ×1Ñ1Ð2CÓDÐà×Ñ Ò!ð .¬c¯j©jÐ9J×9PÑ9PÐQRÑ9SÓ.TÐV_×VfÑVfÐlnÐVfÓVoÐ.oÑp‰Mð .Ü—
‘
Ð,×2Ñ2°1Ñ5Ó6¸Àd×FWÑFWÑ9W×8_Ñ8_ÐegÐ8_Ó8hÐhñˆMñ Ø% }Ð5¸ÈÈÐ8KÑKÐKä-Ø/Ø'Ø)×7Ñ7Ø&×1Ñ1ô	
ð 	
r,   N©TFFT©r$   r%   r&   r   r*   r(   rx   rD   rh   r¹   rw   r+   r,   r-   rí   rí      sZ   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò5ð #Ø"'Ø%*Ø ñ3
ð
 ð3
ð  ð3
ð #ð3
ð ô3
r,   rí   c                   óp   — 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)	ÚFlaxCLIPVisionTransformerrC   rD   c                 ó„  — t        | j                  | j                  ¬«      | _        t	        j
                  | j                  j                  | j                  ¬«      | _        t        | j                  | j                  ¬«      | _	        t	        j
                  | j                  j                  | j                  ¬«      | _
        y rÇ   )rB   rC   rD   rv   r[   rÊ   rË   Úpre_layrnormrç   rï   Úpost_layernormr@   s    r-   rh   zFlaxCLIPVisionTransformer.setupF  sv   € Ü2°4·;±;ÀdÇjÁjÔQˆŒÜŸL™L°·±×1KÑ1KÐSW×S]ÑS]Ô^ˆÔÜ& t§{¡{¸$¿*¹*ÔEˆŒÜ Ÿl™l°4·;±;×3MÑ3MÐUY×U_ÑU_Ô`ˆÕr,   NrŸ   rÝ   c                 ó°  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }| j	                  |«      }| j                  |«      }| j                  |||||¬«      }|d   }|d d …dd d …f   }	| j                  |	«      }	|s
||	f|dd  z   S t        ||	|j                  |j                  ¬«      S )N)rë   rŸ   r    rÜ   rÝ   r   r   rô   )rC   r    rÜ   rö   rv   rÿ   rï   r   r   r"   r#   )
r<   ro   rŸ   r    rÜ   rÝ   r"   rø   r!   rù   s
             r-   rw   z"FlaxCLIPVisionTransformer.__call__L  sþ   € ð 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàŸ™¨Ó5ˆØ×)Ñ)¨-Ó8ˆàŸ,™,Ø'Ø'Ø/Ø!5Ø#ð 'ó 
ˆð ,¨AÑ.ÐØ)ª!¨Q²¨'Ñ2ˆØ×+Ñ+¨MÓ:ˆáØ% }Ð5¸ÈÈÐ8KÑKÐKä-Ø/Ø'Ø)×7Ñ7Ø&×1Ñ1ô	
ð 	
r,   )NTNNT©r$   r%   r&   r   r*   r(   rx   rD   rh   r¹   rw   r+   r,   r-   rý   rý   B  sJ   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òað Ø"ØØ!Ø ñ%
ð ð%
ð ô%
r,   rý   c                   ó8  ‡ — e Zd ZU e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	 	 	 	 	 	 	 	 ddedej                   j"                  dedee   dee   dee   fd„Zˆ xZS )ÚFlaxCLIPTextPreTrainedModelNÚmodule_class©r   r   r   TrC   ÚseedrD   Ú_do_initc                 óZ   •—  | j                   d||dœ|¤Ž}t        ‰| �	  ||||||¬«       y )N©rC   rD   ©Úinput_shaper  rD   r  r+   )r  ÚsuperÚ__init__©	r<   rC   r  r  rD   r  ÚkwargsÚmoduleÚ	__class__s	           €r-   r  z$FlaxCLIPTextPreTrainedModel.__init__x  s=   ø€ ð #�×"Ñ"ÐH¨&¸ÑHÀÑHˆÜ‰Ñ˜ °[ÀtÐSXÐckÐÕlr,   Úrngr  Úparamsr5   c                 óL  — t        j                  |d¬«      }t        j                  t        j                  t        j                  |«      j
                  d   «      |«      }t        j                  |«      }t        j                  j                  |«      \  }}||dœ}	| j                  j                  |	|||«      d   }
|�dt        t        |
«      «      }
t        t        |«      «      }| j                  D ]
  }|
|   ||<   Œ t        «       | _        t!        t#        |«      «      S |
S )NrR   rS   ró   ©r  r�   r  )r(   ÚzerosÚbroadcast_tord   Ú
atleast_2drk   Ú	ones_likerZ   ÚrandomÚsplitr  Úinitr   r	   Ú_missing_keysÚsetr   r   )r<   r  r  r  r�   re   r®   Ú
params_rngr¥   ÚrngsÚrandom_paramsÚmissing_keys               r-   Úinit_weightsz(FlaxCLIPTextPreTrainedModel.init_weights„  sþ   € ä—I‘I˜k°Ô6ˆ	Ü×'Ñ'¬¯
©
´3·>±>À)Ó3L×3RÑ3RÐSUÑ3VÓ(WÐYdÓeˆÜŸ™ yÓ1ˆä"%§*¡*×"2Ñ"2°3Ó"7Ñˆ
�KØ$°Ñ=ˆàŸ™×(Ñ(¨¨y¸.È,ÓWÐX`ÑaˆàÐÜ(¬°-Ó)@ÓAˆMÜ!¤(¨6Ó"2Ó3ˆFØ#×1Ñ1ò A�Ø&3°KÑ&@��{Ò#ðAä!$£ˆDÔÜœ.¨Ó0Ó1Ð1à Ð r,   r¥   Útrainr    rÜ   rÝ   c
                 óp  — |�|n| j                   j                  }|�|n| j                   j                  }|	�|	n| j                   j                  }	|€St	        j
                  t	        j                  t	        j                  |«      j                  d   «      |j                  «      }|€t	        j                  |«      }i }
|�||
d<   | j                  j                  d|xs | j                  it	        j                  |d¬«      t	        j                  |d¬«      t	        j                  |d¬«      | |||	|
¬«	      S )Nró   r�   r  rR   rS   ©r!  )rC   r    rÜ   rÝ   r(   r  rd   r  rk   r  r  Úapplyr  Úarray)r<   r�   r®   re   r  r¥   r%  r    rÜ   rÝ   r!  s              r-   rw   z$FlaxCLIPTextPreTrainedModel.__call__™  s!  € ð 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆàÐÜ×+Ñ+¬C¯J©J´s·~±~ÀiÓ7P×7VÑ7VÐWYÑ7ZÓ,[Ð]f×]lÑ]lÓmˆLàÐ!Ü Ÿ]™]¨9Ó5ˆNð ˆØÐ"Ø)ˆD�‰Oà�{‰{× Ñ Ø�vÒ, §¡Ð-Ü�I‰I�i tÔ,Ü�I‰I�n¨DÔ1Ü�I‰I�l¨$Ô/ØˆIØØ ØØð !ó 

ð 
	
r,   rÃ   ©NNNNFNNN)r$   r%   r&   r   Úconfig_classr  r[   ÚModuler*   r(   rx   ÚintrD   r¹   r  rZ   r  ÚPRNGKeyr   r   r$  Údictr   rw   Ú__classcell__©r  s   @r-   r  r  t  sú   ø… Ø!€LØ"€L�"—)‘)Ó"ð
 ØØŸ;™;Øñ
màð
mð ð	
mð
 �y‰yð
mð õ
mñ! §
¡
× 2Ñ 2ð !Àð !ÐPZð !Ðfpó !ð0 ØØØ*.ØØ,0Ø/3Ø&*ñ'
ð
 ð'
ð —Z‘Z×'Ñ'ð'
ð ð'
ð $ D™>ð'
ð ' t™nð'
ð ˜d‘^÷'
r,   r  c                   óB  ‡ — e Zd ZU eZdZdZej                  e	d<   dde
j                  dfdede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	 	 	 	 	 	 ddedej&                  j(                  dedee   dee   dee   fd„Zˆ xZS )ÚFlaxCLIPVisionPreTrainedModelro   Nr  r   TrC   r  r  rD   r  c                 ó’   •— |€d|j                   |j                   df} | j                  d||dœ|¤Ž}t        ‰| �  ||||||¬«       y )Nr   r   r
  r  r+   )rW   r  r  r  r  s	           €r-   r  z&FlaxCLIPVisionPreTrainedModel.__init__È  s]   ø€ ð ÐØ˜f×/Ñ/°×1BÑ1BÀAÐFˆKØ"�×"Ñ"ÐH¨&¸ÑHÀÑHˆÜ‰Ñ˜ °[ÀtÐSXÐckÐÕlr,   r  r  r5   c                 óž  — t         j                  j                  ||«      }t         j                  j                  |«      \  }}||dœ}| j                  j                  ||«      d   }|�dt        t        |«      «      }t        t        |«      «      }| j                  D ]
  }	||	   ||	<   Œ t        «       | _        t        t        |«      «      S |S )Nr  r  )rZ   r  r]   r  r  r  r   r	   r  r  r   r   )
r<   r  r  r  ro   r   r¥   r!  r"  r#  s
             r-   r$  z*FlaxCLIPVisionPreTrainedModel.init_weightsÖ  sÀ   € ä—z‘z×(Ñ(¨¨kÓ:ˆä"%§*¡*×"2Ñ"2°3Ó"7Ñˆ
�KØ$°Ñ=ˆàŸ™×(Ñ(¨¨|Ó<¸XÑFˆàÐÜ(¬°-Ó)@ÓAˆMÜ!¤(¨6Ó"2Ó3ˆFØ#×1Ñ1ò A�Ø&3°KÑ&@��{Ò#ðAä!$£ˆDÔÜœ.¨Ó0Ó1Ð1à Ð r,   r¥   r%  r    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   rO   r   r   r�   r  rS   r'  )rC   r    rÜ   rÝ   r(   Ú	transposer  r(  r  r)  rx   )	r<   ro   r  r¥   r%  r    rÜ   rÝ   r!  s	            r-   rw   z&FlaxCLIPVisionPreTrainedModel.__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ØØ ØØð !ó 
ð 	
r,   rÃ   )NNFNNN)r$   r%   r&   r   r+  Úmain_input_namer  r[   r,  r*   r(   rx   r   r   r-  rD   r¹   r  rZ   r  r.  r   r$  r/  rw   r0  r1  s   @r-   r3  r3  Ã  s  ø… Ø#€LØ$€OØ"€L�"—)‘)Ó"ð
 (,ØØŸ;™;Øñmà ðmð ˜e‘_ðmð ð	mð
 �y‰yðmð õmñ! §
¡
× 2Ñ 2ð !Àð !ÐPZð !Ðfpó !ð, Ø*.ØØ,0Ø/3Ø&*ñ
ð ð
ð —Z‘Z×'Ñ'ð	
ð
 ð
ð $ D™>ð
ð ' t™nð
ð ˜d‘^÷
r,   r3  c                   óÂ  ‡ — e Zd ZU eZdZej                  ed<   dde	j                  dfdede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	 	 	 	 	 	 	 	 ddedej$                  j&                  dedee   dee   dee   fd„Z	 	 	 	 	 ddedej$                  j&                  fd„Z	 ddedej$                  j&                  fd„Zˆ xZS )ÚFlaxCLIPPreTrainedModelNr  r   TrC   r  r  rD   r  c                 ó¾   •— |€0dd|j                   j                  |j                   j                  dff} | j                  d||dœ|¤Ž}t        ‰| �  ||||||¬«       y )Nr  r   r   r
  r  r+   )Úvision_configrW   r  r  r  r  s	           €r-   r  z FlaxCLIPPreTrainedModel.__init__  so   ø€ ð ÐØ! A v×';Ñ';×'FÑ'FÈ×H\ÑH\×HgÑHgÐijÐ#kÐlˆKØ"�×"Ñ"ÐH¨&¸ÑHÀÑHˆÜ‰Ñ˜ °[ÀtÐSXÐckÐÕlr,   r  r  r5   c                 ó   — t        j                  |d   d¬«      }t        j                  t        j                  t        j                  |«      j
                  d   «      |d   «      }t        j                  |«      }t        j                  j                  ||d   «      }t        j                  j                  |«      \  }}	||	dœ}
| j                  j                  |
||||«      d   }|�dt        t        |«      «      }t        t        |«      «      }| j                  D ]
  }||   ||<   Œ t!        «       | _        t#        t%        |«      «      S |S )Nr   rR   rS   ró   r   r  r  )r(   r  r  rd   r  rk   r  rZ   r  r]   r  r  r  r   r	   r  r  r   r   )r<   r  r  r  r�   re   r®   ro   r   r¥   r!  r"  r#  s                r-   r$  z$FlaxCLIPPreTrainedModel.init_weights  s%  € ä—I‘I˜k¨!™n°DÔ9ˆ	Ü×'Ñ'¬¯
©
´3·>±>À)Ó3L×3RÑ3RÐSUÑ3VÓ(WÐYdÐefÑYgÓhˆÜŸ™ yÓ1ˆä—z‘z×(Ñ(¨¨k¸!©nÓ=ˆä"%§*¡*×"2Ñ"2°3Ó"7Ñˆ
�KØ$°Ñ=ˆàŸ™×(Ñ(¨¨y¸,ÈÐXdÓeÐfnÑoˆàÐÜ(¬°-Ó)@ÓAˆMÜ!¤(¨6Ó"2Ó3ˆFØ#×1Ñ1ò A�Ø&3°KÑ&@��{Ò#ðAä!$£ˆDÔÜœ.¨Ó0Ó1Ð1à Ð r,   r¥   r%  r    rÜ   rÝ   c                 óä  — |�|n| j                   j                  }|	�|	n| j                   j                  }	|
�|
n| j                   j                  }
|€St	        j
                  t	        j                  t	        j                  |«      j                  d   «      |j                  «      }|€t	        j                  |«      }t	        j                  |d«      }i }|�||d<   | j                  j                  d|xs | j                  it	        j                  |d¬«      t	        j                  |t        j                  ¬«      t	        j                  |d¬«      t	        j                  |d¬«      | ||	|
|¬«
      S )Nró   r7  r�   r  rR   rS   r'  )rC   r    rÜ   rÝ   r(   r  rd   r  rk   r  r8  r  r(  r  r)  rx   )r<   r�   ro   r®   re   r  r¥   r%  r    rÜ   rÝ   r!  s               r-   rw   z FlaxCLIPPreTrainedModel.__call__4  sC  € ð 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆàÐÜ×+Ñ+¬C¯J©J´s·~±~ÀiÓ7P×7VÑ7VÐWYÑ7ZÓ,[Ð]f×]lÑ]lÓmˆLàÐ!Ü Ÿ]™]¨9Ó5ˆNä—}‘} \°<Ó@ˆð ˆØÐ"Ø)ˆD�‰Oà�{‰{× Ñ Ø�vÒ, §¡Ð-Ü�I‰I�i tÔ,Ü�I‰I�l¬#¯+©+Ô6Ü�I‰I�n¨DÔ1Ü�I‰I�l¨$Ô/ØˆIØØ ØØð !ó 
ð 	
r,   c           	      óÖ  — |€St        j                  t        j                  t        j                  |«      j                  d   «      |j                  «      }|€t        j
                  |«      }i }|�||d<   d„ }| j                  j                  d|xs | j                  it        j                  |d¬«      t        j                  |d¬«      t        j                  |d¬«      | ||¬«      S )at  
        Args:
            input_ids (`numpy.ndarray` of shape `(batch_size, sequence_length)`):
                Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you
                provide it.

                Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
                [`PreTrainedTokenizer.__call__`] for details.

                [What are input IDs?](../glossary#input-ids)

        Returns:
            text_features (`jnp.ndarray` of shape `(batch_size, output_dim`): The text embeddings obtained by applying
            the projection layer to the pooled output of [`FlaxCLIPTextModel`].

        Examples:

        ```python
        >>> from transformers import AutoTokenizer, FlaxCLIPModel

        >>> model = FlaxCLIPModel.from_pretrained("openai/clip-vit-base-patch32")
        >>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")

        >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="np")
        >>> text_features = model.get_text_features(**inputs)
        ```ró   r�   c                 ó\   — | j                  ||||¬«      }|d   }| j                  |«      }|S )N)r�   r®   re   rŸ   r   )Ú
text_modelÚtext_projection)r  r�   r®   re   rŸ   Útext_outputsrù   Útext_featuress           r-   Ú_get_featuresz@FlaxCLIPPreTrainedModel.get_text_features.<locals>._get_features�  sD   € Ø!×,Ñ,Ø#Ø-Ø)Ø+ð	 -ó ˆLð )¨™OˆMØ"×2Ñ2°=ÓAˆMØ Ð r,   r  rR   rS   ©Úmethodr!  )
r(   r  rd   r  rk   r  r  r(  r  r)  )	r<   r�   r®   re   r  r¥   r%  r!  rF  s	            r-   Úget_text_featuresz)FlaxCLIPPreTrainedModel.get_text_featuresa  sÕ   € ðF ÐÜ×+Ñ+¬C¯J©J´s·~±~ÀiÓ7P×7VÑ7VÐWYÑ7ZÓ,[Ð]f×]lÑ]lÓmˆLàÐ!Ü Ÿ]™]¨9Ó5ˆNð ˆØÐ"Ø)ˆD�‰Oò		!ð �{‰{× Ñ Ø�vÒ, §¡Ð-Ü�I‰I�i tÔ,Ü�I‰I�n¨DÔ1Ü�I‰I�l¨$Ô/ØˆIØ Øð !ó 
ð 	
r,   c                 óî   — t        j                  |d«      }i }|�||d<   d„ }| j                  j                  d|xs | j                  it        j
                  |t         j                  ¬«      | ||¬«      S )a�  
        Args:
            pixel_values (`numpy.ndarray` of shape `(batch_size, num_channels, height, width)`):
                Pixel values. Padding will be ignored by default should you provide it. Pixel values can be obtained
                using [`AutoImageProcessor`]. See [`CLIPImageProcessor.__call__`] for details.

        Returns:
            image_features (`jnp.ndarray` of shape `(batch_size, output_dim`): The image embeddings obtained by
            applying the projection layer to the pooled output of [`FlaxCLIPVisionModel`]

        Examples:

        ```python
        >>> from PIL import Image
        >>> import requests
        >>> from transformers import AutoProcessor, FlaxCLIPModel

        >>> model = FlaxCLIPModel.from_pretrained("openai/clip-vit-base-patch32")
        >>> processor = AutoProcessor.from_pretrained("openai/clip-vit-base-patch32")

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

        >>> inputs = processor(images=image, return_tensors="np")

        >>> image_features = model.get_image_features(**inputs)
        ```r7  r�   c                 óX   — | j                  ||¬«      }|d   }| j                  |«      }|S )N)ro   rŸ   r   )Úvision_modelÚvisual_projection)r  ro   rŸ   Úvision_outputsrù   Úimage_featuress         r-   rF  zAFlaxCLIPPreTrainedModel.get_image_features.<locals>._get_featuresÉ  s8   € Ø#×0Ñ0¸lÐZgÐ0ÓhˆNØ*¨1Ñ-ˆMØ#×5Ñ5°mÓDˆNØ!Ð!r,   r  rS   rG  )r(   r8  r  r(  r  r)  rx   )r<   ro   r  r¥   r%  r!  rF  s          r-   Úget_image_featuresz*FlaxCLIPPreTrainedModel.get_image_features¤  s{   € ô< —}‘} \°<Ó@ˆð ˆØÐ"Ø)ˆD�‰Oò	"ð �{‰{× Ñ Ø�vÒ, §¡Ð-Ü�I‰I�l¬#¯+©+Ô6ØˆIØ Øð !ó 
ð 	
r,   rÃ   r*  )NNNNF)NNF)r$   r%   r&   r   r+  r  r[   r,  r*   r(   rx   r   r   r-  rD   r¹   r  rZ   r  r.  r   r$  r/  rw   rI  rP  r0  r1  s   @r-   r;  r;    sh  ø… Ø€LØ"€L�"—)‘)Ó"ð
 (,ØØŸ;™;Øñmàðmð ˜e‘_ðmð ð	mð
 �y‰yðmð õmñ! §
¡
× 2Ñ 2ð !Àð !ÐPZð !Ðfpó !ð6 ØØØ*.ØØ,0Ø/3Ø&*ñ+
ð ð+
ð —Z‘Z×'Ñ'ð+
ð ð+
ð $ D™>ð+
ð ' t™nð+
ð ˜d‘^ó+
ð` ØØØ*.ØñA
ð
 ðA
ð —Z‘Z×'Ñ'óA
ðH `eñ1
Ø$(ð1
Ø>A¿j¹j×>PÑ>P÷1
r,   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	)ÚFlaxCLIPTextModulerC   rD   c                 óP   — t        | j                  | j                  ¬«      | _        y ré   )rí   rC   rD   rB  r@   s    r-   rh   zFlaxCLIPTextModule.setupÜ  s   € Ü1°$·+±+ÀTÇZÁZÔPˆ�r,   rŸ   r    rÜ   rÝ   c           	      ó2   — | j                  |||||||¬«      S )N©r�   r®   re   rŸ   r    rÜ   rÝ   )rB  )r<   r�   r®   re   rŸ   r    rÜ   rÝ   s           r-   rw   zFlaxCLIPTextModule.__call__ß  s/   € ð �‰ØØ)Ø%Ø'Ø/Ø!5Ø#ð ó 
ð 	
r,   Nrú   rû   r+   r,   r-   rR  rR  Ø  s[   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òQð #Ø"'Ø%*Ø ñ
ð
 ð
ð  ð
ð #ð
ð ô
r,   rR  c                   ó   — e Zd ZeZy)ÚFlaxCLIPTextModelN)r$   r%   r&   rR  r  r+   r,   r-   rW  rW  ô  s   „ Ø%�Lr,   rW  a'  
    Returns:

    Example:

    ```python
    >>> from transformers import AutoTokenizer, FlaxCLIPTextModel

    >>> model = FlaxCLIPTextModel.from_pretrained("openai/clip-vit-base-patch32")
    >>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")

    >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="np")

    >>> outputs = model(**inputs)
    >>> last_hidden_state = outputs.last_hidden_state
    >>> pooler_output = outputs.pooler_output  # pooled (EOS token) states
    ```
)Úoutput_typer+  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	)Ú%FlaxCLIPTextModelWithProjectionModulerC   rD   c                 óÆ   — t        | j                  | j                  ¬«      | _        t	        j
                  | j                  j                  d| j                  ¬«      | _        y )NrS   F)rM   rD   )rí   rC   rD   rB  r[   r�   Úprojection_dimrC  r@   s    r-   rh   z+FlaxCLIPTextModelWithProjectionModule.setup  s>   € Ü1°$·+±+ÀTÇZÁZÔPˆŒÜ!Ÿx™x¨¯©×(BÑ(BÈUÐZ^×ZdÑZdÔeˆÕr,   rŸ   r    rÜ   rÝ   c           	      óÖ   — | j                  |||||||¬«      }|d   }	| j                  |	«      }
|s|
|d   f|dd  z   S t        |
|j                  |j                  |j
                  ¬«      S )NrU  r   r   rO   )r    r!   r"   r#   )rB  rC  r   r!   r"   r#   )r<   r�   r®   re   rŸ   r    rÜ   rÝ   rD  rù   r    s              r-   rw   z.FlaxCLIPTextModelWithProjectionModule.__call__  s�   € ð —‘ØØ)Ø%Ø'Ø/Ø!5Ø#ð 'ó 
ˆð % Q™ˆØ×*Ñ*¨=Ó9ˆáØ ¨a¡Ð1°LÀÀÐ4DÑDÐDä&Ø#Ø*×<Ñ<Ø&×4Ñ4Ø#×.Ñ.ô	
ð 	
r,   Nrú   rû   r+   r,   r-   rZ  rZ    s[   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òfð #Ø"'Ø%*Ø ñ
ð
 ð
ð  ð
ð #ð
ð ô
r,   rZ  c                   ó   — e Zd ZeZy)ÚFlaxCLIPTextModelWithProjectionN)r$   r%   r&   rZ  r  r+   r,   r-   r_  r_  ;  s   „ Ø8�Lr,   r_  aì  
    Returns:

    Example:

    ```python
    >>> from transformers import AutoTokenizer, FlaxCLIPTextModelWithProjection

    >>> model = FlaxCLIPTextModelWithProjection.from_pretrained("openai/clip-vit-base-patch32")
    >>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")

    >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="np")

    >>> outputs = model(**inputs)
    >>> text_embeds = outputs.text_embeds
    ```
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	)ÚFlaxCLIPVisionModulerC   rD   c                 óP   — t        | j                  | j                  ¬«      | _        y ré   )rý   rC   rD   rL  r@   s    r-   rh   zFlaxCLIPVisionModule.setup]  s   € Ü5°d·k±kÈÏÉÔTˆÕr,   rŸ   r    rÜ   rÝ   c                 ó.   — | j                  |||||¬«      S )N©ro   rŸ   r    rÜ   rÝ   )rL  )r<   ro   rŸ   r    rÜ   rÝ   s         r-   rw   zFlaxCLIPVisionModule.__call__`  s+   € ð × Ñ Ø%Ø'Ø/Ø!5Ø#ð !ó 
ð 	
r,   Nrú   r  r+   r,   r-   ra  ra  Y  s[   … ØÓØ—{‘{€Eˆ3�9‰9Ó"òUð #Ø"'Ø%*Ø ñ
ð ð
ð  ð	
ð
 #ð
ð ô
r,   ra  c                   ó   — e Zd ZeZy)ÚFlaxCLIPVisionModelN)r$   r%   r&   ra  r  r+   r,   r-   rf  rf  q  s   „ Ø'�Lr,   rf  a¶  
    Returns:

    Example:

    ```python
    >>> from PIL import Image
    >>> import requests
    >>> from transformers import AutoProcessor, FlaxCLIPVisionModel

    >>> model = FlaxCLIPVisionModel.from_pretrained("openai/clip-vit-base-patch32")
    >>> processor = AutoProcessor.from_pretrained("openai/clip-vit-base-patch32")

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

    >>> inputs = processor(images=image, return_tensors="np")

    >>> outputs = model(**inputs)
    >>> last_hidden_state = outputs.last_hidden_state
    >>> pooler_output = outputs.pooler_output  # pooled CLS states
    ```
c                   ór   — e Zd ZU eed<   ej                  Zej                  ed<   d„ Z	 	 	 	 	 	 	 	 dde	fd„Z
y)ÚFlaxCLIPModulerC   rD   c                 óâ  ‡ — ‰ j                   j                  }‰ j                   j                  }‰ j                   j                  ‰ _        |j                  ‰ _        |j                  ‰ _        t        |‰ j                  ¬«      ‰ _	        t        |‰ j                  ¬«      ‰ _        t        j                  ‰ j                  ‰ j                  t        j                  j                  j!                  d«      d¬«      ‰ _        t        j                  ‰ j                  ‰ j                  t        j                  j                  j!                  d«      d¬«      ‰ _        ‰ j'                  dˆ fd„g «      ‰ _        y )NrS   rG   F)rD   rN   rM   Úlogit_scalec                 ó\   •— t        j                  |«      ‰j                  j                  z  S rÃ   )r(   r—   rC   Úlogit_scale_init_value)Ú_rk   r<   s     €r-   ú<lambda>z&FlaxCLIPModule.setup.<locals>.<lambda>°  s   ø€ ¬C¯H©H°U«O¸d¿k¹k×>`Ñ>`Ñ,`€ r,   )rC   Útext_configr=  r\  rV   Útext_embed_dimÚvision_embed_dimrí   rD   rB  rý   rL  r[   r�   rZ   r\   r]   rM  rC  rY   rj  )r<   ro  r=  s   `  r-   rh   zFlaxCLIPModule.setup—  s
  ø€ Ø—k‘k×-Ñ-ˆØŸ™×1Ñ1ˆà"Ÿk™k×8Ñ8ˆÔØ)×5Ñ5ˆÔØ -× 9Ñ 9ˆÔä1°+ÀTÇZÁZÔPˆŒÜ5°mÈ4Ï:É:ÔVˆÔä!#§¡Ø×ÑØ—*‘*ÜŸ™×+Ñ+×2Ñ2°4Ó8Øô	"
ˆÔô  "Ÿx™xØ×ÑØ—*‘*ÜŸ™×+Ñ+×2Ñ2°4Ó8Øô	 
ˆÔð  Ÿ:™:ØÓ`Ðbdó
ˆÕr,   NrŸ   c	           	      óP  — |�|n| j                   j                  }| j                  |||||¬«      }	| j                  |||||||¬«      }
|	d   }| j	                  |«      }|
d   }| j                  |«      }|t        j                  j                  |dd¬«      z  }|t        j                  j                  |dd¬«      z  }t        j                  | j                  «      }t        j                  ||j                  «      |z  }|j                  }|s|||||
|	fS t        |||||
|	¬«      S )Nrd  rU  r   ró   T)rU   Úkeepdims)r0   r1   r    r2   r3   r4   )rC   rÝ   rL  rB  rM  rC  r(   ÚlinalgÚnormÚexprj  ÚmatmulÚTr/   )r<   r�   ro   r®   re   rŸ   r    rÜ   rÝ   rN  rD  r2   r    rj  r1   r0   s                   r-   rw   zFlaxCLIPModule.__call__³  sM  € ð &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆà×*Ñ*Ø%Ø'Ø/Ø!5Ø#ð +ó 
ˆð —‘ØØ)Ø%Ø'Ø/Ø!5Ø#ð 'ó 
ˆð & aÑ(ˆØ×-Ñ-¨lÓ;ˆà" 1‘oˆØ×*Ñ*¨;Ó7ˆð $¤c§j¡j§o¡o°lÈÐVZ oÓ&[Ñ[ˆØ!¤C§J¡J§O¡O°KÀbÐSW OÓ$XÑXˆô —g‘g˜d×.Ñ.Ó/ˆÜŸ*™* [°,·.±.ÓAÀKÑOˆØ*×,Ñ,ÐáØ$ o°{ÀLÐR^Ð`nÐoÐoäØ-Ø+Ø#Ø%Ø*Ø .ô
ð 	
r,   )NNNNTNNN)r$   r%   r&   r   r*   r(   rx   rD   rh   r¹   rw   r+   r,   r-   rh  rh  “  sH   … ØÓØ—{‘{€Eˆ3�9‰9Ó"ò
ð< ØØØØ"ØØ!Øñ8
ð ô8
r,   rh  c                   ó   — e Zd ZeZy)ÚFlaxCLIPModelN)r$   r%   r&   rh  r  r+   r,   r-   rz  rz  î  s   „ à!�Lr,   rz  ai  
    Returns:

    Example:

    ```python
    >>> import jax
    >>> from PIL import Image
    >>> import requests
    >>> from transformers import AutoProcessor, FlaxCLIPModel

    >>> model = FlaxCLIPModel.from_pretrained("openai/clip-vit-base-patch32")
    >>> processor = AutoProcessor.from_pretrained("openai/clip-vit-base-patch32")

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

    >>> inputs = processor(
    ...     text=["a photo of a cat", "a photo of a dog"], images=image, return_tensors="np", padding=True
    ... )

    >>> outputs = model(**inputs)
    >>> logits_per_image = outputs.logits_per_image  # this is the image-text similarity score
    >>> probs = jax.nn.softmax(logits_per_image, axis=1)  # we can take the softmax to get the label probabilities
    ```
)rz  r;  rW  r  r_  rf  r3  )NÚtypingr   r   r   r   ÚflaxÚ
flax.linenÚlinenr[   rZ   Ú	jax.numpyÚnumpyr(   Úflax.core.frozen_dictr   r   r	   r
   r   Úflax.linen.attentionr   Úflax.traverse_utilr   r   r   Úmodeling_flax_outputsr   r   Úmodeling_flax_utilsr   r   r   r   Úutilsr   r   r   Úconfiguration_clipr   r   r   Ú
get_loggerr$   ÚloggerÚCLIP_START_DOCSTRINGÚCLIP_TEXT_INPUTS_DOCSTRINGÚCLIP_VISION_INPUTS_DOCSTRINGÚCLIP_INPUTS_DOCSTRINGÚstructÚ	dataclassr   r/   r,  rB   rz   r…   r»   rÅ   rÔ   rç   rí   rý   r  r3  r;  rR  rW  ÚFLAX_CLIP_TEXT_MODEL_DOCSTRINGrZ  r_  Ú.FLAX_CLIP_TEXT_MODEL_WITH_PROJECTION_DOCSTRINGra  rf  Ú FLAX_CLIP_VISION_MODEL_DOCSTRINGrh  rz  ÚFLAX_CLIP_MODEL_DOCSTRINGÚ__all__r+   r,   r-   ú<module>r•     sØ  ð÷  /Ó .ã Ý Û 
Ý ß >Ñ >ß 6Ý >ß ;Ý ç X÷ó ÷ @Ñ ?ß LÑ Lð 
ˆ×	Ñ	˜HÓ	%€ð!Ð ðFÐ ð@ Ð ð!Ð ðH ‡�×Ñô9˜kó 9ó ð9ð: ‡�×Ñô 
�[ó  
ó ð 
ôF#˜rŸy™yô #ôL˜RŸY™Yô ô.X˜Ÿ	™	ô Xôv�"—)‘)ô ô('˜2Ÿ9™9ô 'ôT,
˜bŸi™iô ,
ô^
�b—i‘iô 
ô4?
˜bŸi™iô ?
ôD/
 §	¡	ô /
ôdL
Ð"5ô L
ô^E
Ð$7ô E
ôPJ
Ð1ô J
ôZ
˜Ÿ™ô 
ô8&Ð3ô &ð"Ð ñ& Ð*Ð,FÐIgÑ,gÔ hÙ  ØÐ#AÐP^õô
'
¨B¯I©Iô '
ôT9Ð&Aô 9ð2Ð .ñ$ Ø#Ð%?ÐBpÑ%pôñ !Ø#Ð1HÐWeõô

˜2Ÿ9™9ô 
ô0(Ð7ô (ð$Ð  ñ0 Ð,Ð.JÐMmÑ.mÔ nÙ  ØÐ%CÐRbõô
X
�R—Y‘Yô X
ñv Ð*Ó+ô"Ð+ó "ó ,ð"ðÐ ñ6 ˜Ð(=Ð@YÑ(YÔ ZÙ   ¸NÐYcÕ dò�r,   