Ë
    T^(h'  ã                   ó.  — d Z ddlmZmZ ddlZddlmZ ddlm	Z	  G d„ d	ej                  j                  j                  «      Z G d
„ dej                  j                  j                  «      Z G d„ dej                  j                  j                  «      Zy)a  

Generic interface to various configurations of the Perceiver Resampler, that simply takes in a series of (potentially
time-indexed) contextual embeddings, and "resamples" (compresses) them down to a pre-specified number of latents! Note
that the Perceiver in general resamples based solely off the *long-range* context; there's a nice opportunity here to
prime the Perceiver Resampler with say a single layer's worth of language embeddings (the target domain), and use that
to softly "retrieve & compress" what we need --> this would be a novel contribution we should explore.

References:
    - DeepMind's Flamingo: https://www.deepmind.com/blog/tackling-multiple-tasks-with-a-single-visual-language-model
    - Code borrowed w/ love from: https://github.com/lucidrains/flamingo-pytorch

é    )ÚOptionalÚTupleNé   )Ú
shape_listé   )ÚIdeficsConfigc                   ó~   ‡ — e Zd Zdededededededdfˆ fd	„Zˆ fd
„Zdej                  dej                  fd„Z	ˆ xZ
S )ÚTFIdeficsPerceiverResamplerÚconfigÚ	embed_dimÚdepthÚn_headsÚhead_dimÚ	n_latentsÚreturnNc                 óŽ  •— t        ‰	| �  di |¤Ž ||||f\  | _        | _        | _        | _        |j                  j                  | _        t        |j                  d«      s| j                  dz  n|j                  j                  dz  | _        g | _        t        |«      D ]s  }| j                  j                  t        | j                  | j                  | j                  | j                  d|› d�¬«      t!        | j                  |d|› d�¬«      g«       Œu t"        j$                  j&                  j)                  dd¬	«      | _        y
)ao  
        Instantiates a Perceiver Resampler that operates over a sequence of embeddings (say from a ResNet or ViT or
        MAE) of a given dimension, performs `depth` blocks of cross-attention with a fixed `n_latents` inputs, then
        returns a Tensor of shape [bsz, n_latents, embed_dim]. :param embed_dim: Dimensionality of embeddings being fed
        to the Perceiver Resampler (also dimensionality of latent embeddings *returned* by the Perceiver Resampler.
        Could be e.g., VIT embed_dim, ResNet pool dim, and so on.

        Args:
            config (`IdeficsConfig`): config object
            embed_dim (`int`): The size of each embedding vector
            depth (`int`): Depth of the Perceiver Resampler (Transformer w/ cross attention). Should be shallow (< 3).
            n_heads (`int`): Number of heads in each Transformer block (for multi-headed self-attention).
            head_dim (`int`): Dimensionality of each head projection in the Transformer block.
            n_latents (`int`):
                Number of latent embeddings to resample ("compress") the input sequence to (usually < 128).

        r   é   zblocks.z.0©Únamez.1çñhãˆµøä>Ú
layer_norm©Úepsilonr   N© )ÚsuperÚ__init__r   r   r   r   Úperceiver_configÚqk_layer_norms_perceiverÚqk_layer_normsÚhasattrÚvision_configÚintermediate_dimÚblocksÚrangeÚappendÚTFIdeficsPerceiverAttentionÚTFIdeficsMLPÚtfÚkerasÚlayersÚLayerNormalizationr   )
Úselfr   r   r   r   r   r   ÚkwargsÚiÚ	__class__s
            €úf/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/idefics/perceiver_tf.pyr   z$TFIdeficsPerceiverResampler.__init__1   s+  ø€ ô( 	‰ÑÑ"˜6Ò"ØFOÐQXÐZbÐdmÐFmÑCˆŒ˜œ d¤m°T´^Ø$×5Ñ5×NÑNˆÔô ˜6×/Ñ/°Ô=ð �N‰N˜QÒà×%Ñ%×/Ñ/°!Ñ3ð 	Ôð ˆŒÜ�u“ò 	ˆAØ�K‰K×Ñä/ØŸ™¨¯©°d·m±mÀT×EXÑEXÐahÐijÐhkÐkmÐ_nôô ! ×!6Ñ!6¸ÀwÈqÈcÐQSÀ_ÔUð	õð	ô Ÿ(™(Ÿ/™/×<Ñ<ÀTÐP\Ð<Ó]ˆ�ó    c                 ó„   •— | j                  | j                  | j                  fddd¬«      | _        t        ‰| �  |«       y )NÚrandom_normalTÚlatents)ÚshapeÚinitializerÚ	trainabler   )Ú
add_weightr   r   r4   r   Úbuild)r,   Úinput_shaper/   s     €r0   r9   z!TFIdeficsPerceiverResampler.build\   s>   ø€ à—‘Ø—>‘> 4§>¡>Ð2ÀÐ[_Ðfoð 'ó 
ˆŒô 	‰‰�kÕ"r1   Úcontextc                 ó  — t        j                  | j                  d¬«      }t        j                  |t        j                  |«      d   ddg«      }| j
                  D ]  \  }} |||«      |z   } ||«      |z   }Œ | j                  |«      S )zWResample arbitrary length context & *compress* down to self.n_latents latent embeddingsr   ©Úaxisr   )r(   Úexpand_dimsr4   Útiler5   r#   r   )r,   r;   r4   ÚattnÚffs        r0   Úcallz TFIdeficsPerceiverResampler.callc   s„   € ô —.‘. §¡°AÔ6ˆÜ—'‘'˜'¤B§H¡H¨WÓ$5°aÑ$8¸!¸QÐ#?Ó@ˆàŸ™ò 	,‰HˆD�"Ù˜7 GÓ,¨wÑ6ˆGÙ˜“k GÑ+‰Gð	,ð �‰˜wÓ'Ð'r1   )Ú__name__Ú
__module__Ú__qualname__r   Úintr   r9   r(   ÚTensorrC   Ú__classcell__©r/   s   @r0   r
   r
   0   si   ø„ ð)^Ø#ð)^Ø03ð)^Ø<?ð)^ØJMð)^ØY\ð)^Øilð)^à	õ)^ôV#ð	(˜BŸI™Ið 	(¨"¯)©)÷ 	(r1   r
   c            
       ó„   ‡ — e Zd Zdededededdf
ˆ fd„Zdej                  d	ej                  dej                  fd
„Zˆ xZ	S )r&   r   r   r   r   r   Nc                 ó0  •— t        ‰| �  di |¤Ž |||c| _        | _        | _        || _        t        j                  j                  j                  dd¬«      | _
        t        j                  j                  j                  dd¬«      | _        | j
                  r`t        j                  j                  j                  dd¬«      | _        t        j                  j                  j                  dd¬«      | _        | j                  dz  | _        t        j                  j                  j                  | j                  | j                  z  dd	¬
«      | _        t        j                  j                  j                  | j                  | j                  z  dd¬
«      | _        t        j                  j                  j                  | j                  | j                  z  dd¬
«      | _        t        j                  j                  j                  |dd¬
«      | _        y)ziPerceiver Cross-Attention Module --> let long-form inputs be `context`, resampled embeddings be `latents`r   Úcontext_layer_normr   Úlatents_layer_normÚq_layer_normÚk_layer_normg      à¿FÚq_proj©Úuse_biasr   Úk_projÚv_projÚoutput_projNr   )r   r   r   r   r   r   r(   r)   r*   r+   rM   rN   rO   rP   Úqk_scaleÚDenserQ   rT   rU   rV   )r,   r   r   r   r   r-   r/   s         €r0   r   z$TFIdeficsPerceiverAttention.__init__p   sƒ  ø€ ä‰ÑÑ"˜6Ò"Ø6?ÀÈ(Ð3ˆŒ˜œ d¤mØ,ˆÔä"$§(¡(§/¡/×"DÑ"DÈTÐXlÐ"DÓ"mˆÔÜ"$§(¡(§/¡/×"DÑ"DÈTÐXlÐ"DÓ"mˆÔØ×ÒÜ "§¡§¡× BÑ BÈ4ÐVdÐ BÓ eˆDÔÜ "§¡§¡× BÑ BÈ4ÐVdÐ BÓ eˆDÔàŸ™ tÑ+ˆŒô —h‘h—o‘o×+Ñ+¨D¯L©L¸4¿=¹=Ñ,HÐSXÐ_gÐ+ÓhˆŒÜ—h‘h—o‘o×+Ñ+¨D¯L©L¸4¿=¹=Ñ,HÐSXÐ_gÐ+ÓhˆŒÜ—h‘h—o‘o×+Ñ+¨D¯L©L¸4¿=¹=Ñ,HÐSXÐ_gÐ+ÓhˆŒäŸ8™8Ÿ?™?×0Ñ0°ÀUÐQ^Ð0Ó_ˆÕr1   r;   r4   c                 óô  — | j                  |«      }| j                  |«      }t        |«      \  }}}| j                  |«      }| j	                  t        j                  ||gd¬«      «      }| j                  t        j                  ||gd¬«      «      }|||fD �	cg c]T  }	t        j                  t        j                  |	||	j                  d   | j                  | j                  f«      g d¢¬«      ‘ŒV c}	\  }}}| j                  r"| j                  |«      }| j                  |«      }t        j                   d|| j"                  z  |«      }
|
t        j$                  |
dd¬	«      z
  }t
        j&                  j)                  |d¬«      }t        j                   d
||«      }| j+                  t        j                  t        j                  |g d¢¬«      |d| j                  | j                  z  f«      «      S c c}	w )a=  
        Runs Perceiver Self-Attention, with special (context, latents) appended along the `seq` dimension!

        Args:
            context (`tf.Tensor`):
                Tensor of shape `[bsz, seq, embed_dim]` representing long-form context to resample.
            latents (`tf.Tensor`):
                Tensor of shape `[bsz, n_latents, embed_dim]` representing fixed length latents to compress to.

        Returns:
            `tf.Tensor`: Tensor of shape `[bsz, n_latents, embed_dim]` representing attention over latents w/ cross
            from context.
        éþÿÿÿr=   r   )r   é   r   r   )Úpermz... i d, ... j d -> ... i jéÿÿÿÿT)r>   Úkeepdimsz... i j, ... j d -> ... i d)rM   rN   r   rQ   rT   r(   ÚconcatrU   Ú	transposeÚreshaper5   r   r   r   rO   rP   ÚeinsumrW   Ú
reduce_maxÚnnÚsoftmaxrV   )r,   r;   r4   Ú
batch_sizeÚ
seq_lengthr   ÚqÚkÚvÚxÚscoresÚstabilized_scoresrA   Ú	resampleds                 r0   rC   z TFIdeficsPerceiverAttention.call…   s®  € ð ×)Ñ)¨'Ó2ˆØ×)Ñ)¨'Ó2ˆÜ,6°wÓ,?Ñ)ˆ
�J 	ð �K‰K˜Ó ˆØ�K‰KœŸ	™	 7¨GÐ"4¸2Ô>Ó?ˆØ�K‰KœŸ	™	 7¨GÐ"4¸2Ô>Ó?ˆð ˜˜A�Yö
àô �L‰LœŸ™ A¨
°A·G±G¸A±JÀÇÁÈdÏmÉmÐ'\Ó]ÒdpÖqò
‰ˆˆ1ˆað
 ×ÒØ×!Ñ! !Ó$ˆAØ×!Ñ! !Ó$ˆAä—‘Ð8¸!¸d¿m¹mÑ:KÈQÓOˆØ"¤R§]¡]°6ÀÈTÔ%RÑRÐÜ�u‰u�}‰}Ð.°Rˆ}Ó8ˆô —I‘IÐ;¸TÀ1ÓEˆ	Ø×ÑÜ�J‰J”r—|‘| I²LÔAÀJÐPRÐTX×T`ÑT`Ðcg×cpÑcpÑTpÐCqÓró
ð 	
ùò
s   ÂAG5)
rD   rE   rF   rG   Úboolr   r(   rH   rC   rI   rJ   s   @r0   r&   r&   o   sY   ø„ ð` #ð `°ð `¸sð `ÐTXð `Ðgkõ `ð*+
˜BŸI™Ið +
°·	±	ð +
¸b¿i¹i÷ +
r1   r&   c                   óh   ‡ — e Zd Zdefˆ fd„Zdeeej                        dej                  fd„Z	ˆ xZ
S )r'   r   c                 óð  •— t        ‰| �  di |¤Ž |j                  j                  | _        t        j
                  j                  j                  dd¬«      | _        t        j
                  j                  j                  |dd¬«      | _
        t        j
                  j                  j                  d¬«      | _        t        j
                  j                  j                  | j                  dd	¬«      | _        y
)z:Simple MLP block with intermediate_size and embedding sizer   Úlnr   FÚfcrR   Úactr   Úc_projNr   )r   r   r!   r   r(   r)   r*   r+   rr   rX   rs   ÚReLUrt   ru   )r,   Úintermediate_sizer   r-   r/   s       €r0   r   zTFIdeficsMLP.__init__´   s«   ø€ ä‰ÑÑ"˜6Ò"Ø×-Ñ-×7Ñ7ˆŒÜ—(‘(—/‘/×4Ñ4¸TÈÐ4ÓMˆŒÜ—(‘(—/‘/×'Ñ'Ð(9ÀEÐPTÐ'ÓUˆŒÜ—8‘8—?‘?×'Ñ'¨UÐ'Ó3ˆŒÜ—h‘h—o‘o×+Ñ+¨D¯N©NÀUÐQYÐ+ÓZˆ�r1   Úhidden_statesr   c                 óŽ   — | j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }|S )N)rr   rs   rt   ru   )r,   rx   s     r0   rC   zTFIdeficsMLP.call½   s@   € ØŸ™ Ó.ˆØŸ™ Ó.ˆØŸ™ Ó/ˆØŸ™ MÓ2ˆàÐr1   )rD   rE   rF   r   r   r   r   r(   rH   rC   rI   rJ   s   @r0   r'   r'   ³   s6   ø„ ð[°-õ [ð (¨5°·±Ñ+;Ñ"<ð ÀÇÁ÷ r1   r'   )Ú__doc__Útypingr   r   Ú
tensorflowr(   Úmodeling_tf_utilsr   Úconfiguration_ideficsr   r)   r*   ÚLayerr
   r&   r'   r   r1   r0   ú<module>r€      sj   ðñ4÷ #ã å +Ý 0ô<( "§(¡(§/¡/×"7Ñ"7ô <(ô~A
 "§(¡(§/¡/×"7Ñ"7ô A
ôH�2—8‘8—?‘?×(Ñ(õ r1   