§
    ‚ŠtjÜ[ ã                   ó  — d Z ddlZddlZddlmZ ddlmZ ddlmZ ddlZddlm	Z	 ddl
mZ dd	lmZ dd
lmZ ddlmZ ddlmZmZ ddlmZ ddlmZ ddlmZmZmZmZmZm Z  ddl!m"Z"m#Z#m$Z$m%Z%m&Z&  ej'        e(¦  «        Z)dZ*dZ+dZ,e&e$z  e%z  Z- ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z. ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z/ ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z0 G d„ de	j1        ¦  «        Z2 G d „ d!e	j1        ¦  «        Z3 G d"„ d#e	j1        ¦  «        Z4 G d$„ d%e	j1        ¦  «        Z5 G d&„ d'e	j1        ¦  «        Z6 G d(„ d)e	j1        ¦  «        Z7 G d*„ d+e	j1        ¦  «        Z8 G d,„ d-e	j1        ¦  «        Z9 G d.„ d/e¦  «        Z: G d0„ d1e	j1        ¦  «        Z; G d2„ d3e	j1        ¦  «        Z<e G d4„ d5e¦  «        ¦   «         Z=e G d6„ d7e=¦  «        ¦   «         Z>e G d8„ d9e=¦  «        ¦   «         Z?e G d:„ d;e=¦  «        ¦   «         Z@e G d<„ d=e=¦  «        ¦   «         ZA G d>„ d?e	j1        ¦  «        ZB G d@„ dAe	j1        ¦  «        ZC G dB„ dCe	j1        ¦  «        ZD edD¬¦  «         G dE„ dFe=¦  «        ¦   «         ZE G dG„ dHe	j1        ¦  «        ZF G dI„ dJe	j1        ¦  «        ZG G dK„ dLe	j1        ¦  «        ZH G dM„ dNe	j1        ¦  «        ZI edO¬¦  «         G dP„ dQe=¦  «        ¦   «         ZJg dR¢ZKdS )SzPyTorch FLAVA model.é    N)ÚOrderedDict)Ú	dataclass)ÚAny)Únné   )Úinitialization)ÚACT2FN)Úcreate_bidirectional_mask)ÚGradientCheckpointingLayer)ÚBaseModelOutputÚBaseModelOutputWithPooling)ÚPreTrainedModel)ÚUnpack)ÚModelOutputÚTransformersKwargsÚauto_docstringÚcan_return_tupleÚloggingÚ	torch_inté   )ÚFlavaConfigÚFlavaImageCodebookConfigÚFlavaImageConfigÚFlavaMultimodalConfigÚFlavaTextConfigzfacebook/flava-image-codebookg$(~Œ¹k@a–  
    Output from FlavaModel containing embeddings and outputs from individual encoders.

    Note that `image_embeddings` and `text_embeddigns` returned are similar to pooled output returned from a
    transformer. If you want embeddings for contrastive loss or retrieval use a FLAVA model's `image_projection` and
    `text_projection` layers on `image_embeddings` and `text_embeddings` respectively.
    )Úcustom_introc                   óÂ   — e Zd ZU dZdZej        dz  ed<   dZe	dz  ed<   dZ
ej        dz  ed<   dZe	dz  ed<   dZej        dz  ed<   dZe	dz  ed<   d	ee         fd
„ZdS )ÚFlavaModelOutputaì  
    image_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `pixel_values` are present):
        The image embeddings which are basically the pooled output of [`FlavaImageModel`].
    image_output (`BaseModelOutputWithPooling`, *optional*, returned when `pixel_values` are present):
        The output of the [`FlavaImageModel`].
    text_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `input_ids` are present):
        The text embeddings which are basically the pooled output of [`FlavaTextModel`].
    text_output (`BaseModelOutputWithPooling`, *optional*, returned when `input_ids` are present):
        The output of the [`FlavaTextModel`].
    multimodal_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `input_ids` and `pixel_values` are present and `skip_multimodal_encoder` is `None` or `False`):
        The multimodal embeddings which are basically the pooled output of [`FlavaTextModel`].
    multimodal_output (`BaseModelOutputWithPooling`, returned when `input_ids` and `pixel_values` are present and `skip_multimodal_encoder` is `None` or `False`):
        The output of the [`FlavaMultimodalModel`].
    NÚimage_embeddingsÚimage_outputÚtext_embeddingsÚtext_outputÚmultimodal_embeddingsÚmultimodal_outputÚreturnc                 ó^   ‡ — t          ˆ fd„‰                      ¦   «         D ¦   «         ¦  «        S )Nc              3   ót   •K  — | ]2}|d vr‰|         n!t          ‰|¦  «                             ¦   «         V — Œ3dS ))r"   r    r$   N©ÚgetattrÚto_tuple)Ú.0ÚkÚselfs     €úf/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/flava/modeling_flava.pyú	<genexpr>z,FlavaModelOutput.to_tuple.<locals>.<genexpr>V   sc   øè è € ð 
ð 
àð Ð TÐTÐTˆD�ŒGˆGÕZaÐbfÐhiÑZjÔZj×ZsÒZsÑZuÔZuð
ð 
ð 
ð 
ð 
ð 
ó    ©ÚtupleÚkeys©r-   s   `r.   r*   zFlavaModelOutput.to_tupleU   sC   ø€ Ýð 
ð 
ð 
ð 
à—Y’Y‘[”[ð
ñ 
ô 
ñ 
ô 
ð 	
r0   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚFloatTensorÚ__annotations__r    r   r!   r"   r#   r$   r2   r   r*   © r0   r.   r   r   4   sÎ   € € € € € € ðð ð 26Ð�eÔ'¨$Ñ.Ð5Ð5Ñ5Ø6:€LÐ,¨tÑ3Ð:Ð:Ñ:Ø04€O�UÔ&¨Ñ-Ð4Ð4Ñ4Ø59€KÐ+¨dÑ2Ð9Ð9Ñ9Ø6:Ð˜5Ô,¨tÑ3Ð:Ð:Ñ:Ø;?ÐÐ1°DÑ8Ð?Ð?Ñ?ð
˜% œ*ð 
ð 
ð 
ð 
ð 
ð 
r0   r   z@
    Class representing pretraining losses from FLAVA model
    c                   óÔ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
ej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   d	efd
„ZdS )ÚFlavaLossesa¯  
    mim (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `mim_labels` and `pixel_values` are present, `input_ids_masked` is absent and `mim_weight` > 0.):
        Masked Image Modeling loss as used in BeIT calculated only for unimodal image data.
    mlm (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `mlm_labels` and `input_ids_masked` are present, `pixel_values` is absent and `mlm_weight` > 0.):
        Masked Language Modeling loss as used in BERT calculated only for unimodal text data.
    itm (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `itm_labels`, `input_ids_masked`, `pixel_values` are present and `itm_weight` > 0.):
        Image Text Matching (ITM) loss calculated for paired image-text data. Note that ITM loss is calculated on
        masked pairs in FLAVA.
    global_contrastive (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `input_ids` and `pixel_values` are present and `global_contrastive_weight` > 0.):
        Contrastive loss for image-text similarity similar to CLIP but calculated globally for paired image-text
        data. This is calculated on unmasked images and texts.
    mmm_image (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `mim_labels`, `pixel_values` and `input_ids_masked` are present and `mmm_image_weight` > 0.):
        Masked Multimodal Modeling loss's image component calculated on paired image-text data.
    mmm_text (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `mlm_labels`, `pixel_values` and `input_ids_masked` are present and `mmm_text_weight` > 0.):
        Masked Multimodal Modeling loss's text component calculated on paired image-text data.
    NÚmimÚmlmÚitmÚglobal_contrastiveÚ	mmm_imageÚmmm_textr%   c                 óD   — d}|                       ¦   «         D ]}|�d} nŒ	|S )NTF)Úvalues)r-   Úall_noneÚvs      r.   rG   zFlavaLosses.all_none{   s9   € ØˆØ—’‘”ð 	ð 	ˆAØˆ}Ø �Ø�ð ð ˆr0   )r5   r6   r7   r8   r?   r9   r:   r;   r@   rA   rB   rC   rD   ÚboolrG   r<   r0   r.   r>   r>   \   sÎ   € € € € € € ðð ð" %)€CˆÔ	˜TÑ	!Ð(Ð(Ñ(Ø$(€CˆÔ	˜TÑ	!Ð(Ð(Ñ(Ø$(€CˆÔ	˜TÑ	!Ð(Ð(Ñ(Ø37Ð˜Ô)¨DÑ0Ð7Ð7Ñ7Ø*.€IˆuÔ  4Ñ'Ð.Ð.Ñ.Ø)-€HˆeÔ $Ñ&Ð-Ð-Ñ-ð˜$ð ð ð ð ð ð r0   r>   a   
    Output from FlavaForPreTraining containing embeddings, and outputs from individual encoders.

    Note that `image_embeddings` and `text_embeddings` returned are similar to pooled output returned from a
    transformer. If you want embeddings for contrastive loss or retrieval use a FLAVA model's `image_projection` and
    `text_projection` layers on `image_embeddings` and `text_embeddings` respectively.
    c                   óV  — e Zd ZU dZdZej        dz  ed<   dZe	ed<   dZ
ej        dz  ed<   dZedz  ed<   dZej        dz  ed<   dZedz  ed<   dZej        dz  ed	<   dZedz  ed
<   dZej        dz  ed<   dZedz  ed<   dZej        dz  ed<   dZedz  ed<   dZej        dz  ed<   dZedz  ed<   dZej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dee         fd„Z dS )ÚFlavaForPreTrainingOutputay  
    loss (`torch.FloatTensor`, *optional*, returned when `return_loss` is True):
        Total loss calculated for this model.
    loss_info (`FlavaLosses`):
        Detailed info for FLAVA Pretraining losses. Check `FlavaLosses` class description for the information on
        the keys.
    image_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `pixel_values` are present):
        The image embeddings which are basically the pooled output of [`FlavaImageModel`].
    image_output (`BaseModelOutputWithPooling`, *optional*, returned when `pixel_values` are present):
        The output of the [`FlavaImageModel`].
    text_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `input_ids` are present):
        The text embeddings which are basically the pooled output of [`FlavaTextModel`].
    text_output (`BaseModelOutputWithPooling`, *optional*, returned when `input_ids` are present):
        The output of the [`FlavaTextModel`].
    multimodal_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `input_ids` and `pixel_values` are present and `skip_unmasked_multimodal_encoder` is `None` or `False`):
        The multimodal embeddings which are basically the pooled output of [`FlavaTextModel`].
    multimodal_output (`BaseModelOutputWithPooling`, returned when `input_ids` and `pixel_values` are present and `skip_unmasked_multimodal_encoder` is `None` or `False`):
        The output of the [`FlavaMultimodalModel`].
    image_masked_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `pixel_values` are present):
        The image embeddings which are basically the pooled output of [`FlavaImageModel`]. Uses `bool_masked_pos`
        to create masked images.
    image_masked_output (`BaseModelOutputWithPooling`, *optional*, returned when `pixel_values` are present):
        The output of the [`FlavaImageModel`]. Uses `bool_masked_pos` to create masked images.
    text_masked_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `input_ids_masked` are present):
        The text embeddings which are basically the pooled output of [`FlavaTextModel`].
    text_masked_output (`BaseModelOutputWithPooling`, *optional*, returned when `input_ids_masked` are present):
        The output of the [`FlavaTextModel`].
    multimodal_masked_embeddings (`torch.FloatTensor` of shape `(batch_size, output_dim)`, *optional*, returned when `input_ids` and `pixel_values` are present):
        The multimodal embeddings which are basically the pooled output of [`FlavaTextModel`].
    multimodal_masked_output (`BaseModelOutputWithPooling`, *optional*, returned when `input_ids_masked` and `pixel_values` are present):
        The output of the [`FlavaMultimodalModel`].
    mim_logits (`torch.FloatTensor` of shape `(batch_size, num_image_patches, image_vocab_size)` or of shape `(total_masked_patches, image_vocab_size)` , *optional*, returned when `pixel_values` are present and `input_ids_masked` are not):
        The logits for MIM unimodal loss. Uses `book_masked_pos` to get masked patches. The flattened output is
            returned when `bool_masked_pos` has some of the patches masked.
    mlm_logits (`torch.FloatTensor` of shape `(batch_size, text_seq_length, text_vocab_size)` or of shape `(total_masked_seq_length, text_vocab_size)`, *optional*, returned when `input_ids_masked` are present and `pixel_values` are not):
        The logits for MLM unimodal loss. The flattened output is returned when `input_ids_masked` has some of
            the tokens masked.
    itm_logits (`torch.FloatTensor` of shape `(batch_size, 2)`, *optional*, returned when `input_ids_masked` and `pixel_values` are present):
        The logits for ITM loss. Note that ITM loss is calculated on masked pairs in FLAVA.
    contrastive_logits_per_image (`torch.FloatTensor` of shape `(image_batch_size, text_batch_size)`):
        The scaled dot product scores between `image_embeddings` and `text_embeddings` but passed through FLAVA's
        `image_projection` and `text_projection` layers respectively. This represents the image-text similarity
        scores. This is calculated on unmasked images and texts.
    contrastive_logits_per_text (`torch.FloatTensor` of shape `(text_batch_size, image_batch_size)`):
        The scaled dot product scores between `text_embeddings` and `image_embeddings` but passed through FLAVA's
        `text_projection` and `image_projection` layers respectively. This is calculated on unmasked images and
        texts.
    mmm_image_logits (`torch.FloatTensor` of shape `(batch_size, num_image_patches, image_vocab_size)` or of shape`(total_masked_patches, image_vocab_size)`, *optional*, returned when `pixel_values` and `input_ids_masked` are present):
        The logits for MMM image multimodal loss. Uses `book_masked_pos` to get masked patches. The flattened
            output is returned when `bool_masked_pos` has some of the patches masked.
    mmm_text_logits (`torch.FloatTensor` of shape `(batch_size, text_seq_length, text_vocab_size)` or of shape `(`(total_masked_seq_length, text_vocab_size)`), *optional*, returned when `pixel_values` and `input_ids_masked` are present):
        The logits for MMM text multimodal loss. The flattened output is returned when `input_ids_masked` has
            some of the tokens masked.
    NÚlossÚ	loss_infor   r    r!   r"   r#   r$   Úimage_masked_embeddingsÚimage_masked_outputÚtext_masked_embeddingsÚtext_masked_outputÚmultimodal_masked_embeddingsÚmultimodal_masked_outputÚ
mim_logitsÚ
mlm_logitsÚ
itm_logitsÚcontrastive_logits_per_imageÚcontrastive_logits_per_textÚmmm_image_logitsÚmmm_text_logitsr%   c                 ój   ‡ ‡— g d¢Št          ˆ ˆfd„‰                      ¦   «         D ¦   «         ¦  «        S )N)r"   r    r$   rQ   rO   rS   c              3   ót   •K  — | ]2}|‰vr‰|         n!t          ‰|¦  «                             ¦   «         V — Œ3d S ©Nr(   )r+   r,   r-   Útransformer_outputss     €€r.   r/   z5FlavaForPreTrainingOutput.to_tuple.<locals>.<genexpr>å   sN   øè è € ÐsÐsÐbc Ð)<Ð <Ð <�T˜!”W�WÅ'È$ÐPQÑBRÔBR×B[ÒB[ÑB]ÔB]ÐsÐsÐsÐsÐsÐsr0   r1   )r-   r^   s   `@r.   r*   z"FlavaForPreTrainingOutput.to_tupleÜ   sK   øø€ ð
ð 
ð 
Ðõ ÐsÐsÐsÐsÐsÐgk×gpÒgpÑgrÔgrÐsÑsÔsÑsÔsÐsr0   )!r5   r6   r7   r8   rL   r9   r:   r;   rM   r>   r   r    r   r!   r"   r#   r$   rN   rO   rP   rQ   rR   rS   rT   rU   rV   rW   rX   rY   rZ   r2   r   r*   r<   r0   r.   rK   rK   „   s9  € € € € € € ð5ð 5ðn &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø!€Iˆ{Ð!Ð!Ñ!Ø15Ð�eÔ'¨$Ñ.Ð5Ð5Ñ5Ø6:€LÐ,¨tÑ3Ð:Ð:Ñ:Ø04€O�UÔ&¨Ñ-Ð4Ð4Ñ4Ø59€KÐ+¨dÑ2Ð9Ð9Ñ9Ø6:Ð˜5Ô,¨tÑ3Ð:Ð:Ñ:Ø;?ÐÐ1°DÑ8Ð?Ð?Ñ?Ø8<Ð˜UÔ.°Ñ5Ð<Ð<Ñ<Ø=AÐÐ3°dÑ:ÐAÐAÑAØ7;Ð˜EÔ-°Ñ4Ð;Ð;Ñ;Ø<@ÐÐ2°TÑ9Ð@Ð@Ñ@Ø=AÐ  %Ô"3°dÑ":ÐAÐAÑAØBFÐÐ8¸4Ñ?ÐFÐFÑFØ+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø=AÐ  %Ô"3°dÑ":ÐAÐAÑAØ<@Ð Ô!2°TÑ!9Ð@Ð@Ñ@Ø15Ð�eÔ'¨$Ñ.Ð5Ð5Ñ5Ø04€O�UÔ&¨Ñ-Ð4Ð4Ñ4ð	t˜% œ*ð 	tð 	tð 	tð 	tð 	tð 	tr0   rK   c            	       ó    ‡ — e Zd ZdZddededdfˆ fd„Zdej        d	e	d
e	dej        fd„Z
	 	 ddej        dej        dz  dedej        fd„Zˆ xZS )ÚFlavaImageEmbeddingszb
    Construct the CLS token, position and patch embeddings. Optionally, also the mask token.
    FÚconfigÚuse_mask_tokenr%   Nc                 óf  •— t          ¦   «                              ¦   «          |p|j        }t          j        t          j        dd|j        ¦  «        ¦  «        | _        |r-t          j        t          j        dd|j        ¦  «        ¦  «        nd | _        t          |j
        |j        |j        |j        ¬¦  «        | _        | j        j        }t          j        t          j        d|dz   |j        ¦  «        ¦  «        | _        t          j        |j        ¦  «        | _        |j        | _        || _        d S )Nr   )Ú
image_sizeÚ
patch_sizeÚnum_channelsÚ	embed_dim)ÚsuperÚ__init__Ú
mask_tokenr   Ú	Parameterr9   ÚzerosÚhidden_sizeÚ	cls_tokenÚPatchEmbeddingsrd   re   rf   Úpatch_embeddingsÚnum_patchesÚposition_embeddingsÚDropoutÚhidden_dropout_probÚdropoutra   )r-   ra   rb   rq   Ú	__class__s       €r.   ri   zFlavaImageEmbeddings.__init__ï   s  ø€ Ý‰Œ×ÒÑÔÐà'Ð<¨6Ô+<ˆÝœ¥e¤k°!°Q¸Ô8JÑ&KÔ&KÑLÔLˆŒØQ_Ði�"œ,¥u¤{°1°a¸Ô9KÑ'LÔ'LÑMÔMÐMÐeiˆŒÝ /ØÔ(ØÔ(ØÔ,ØÔ(ð	!
ñ !
ô !
ˆÔð Ô+Ô7ˆÝ#%¤<µ´¸A¸{ÈQ¹ÐPVÔPbÑ0cÔ0cÑ#dÔ#dˆÔ Ý”z &Ô"<Ñ=Ô=ˆŒØ Ô+ˆŒØˆŒˆˆr0   Ú
embeddingsÚheightÚwidthc                 ó”  — |j         d         dz
  }| j        j         d         dz
  }t          j                             ¦   «         s||k    r||k    r| j        S | j        dd…dd…f         }| j        dd…dd…f         }|j         d         }|| j        z  }	|| j        z  }
t          |dz  ¦  «        }|                     d|||¦  «        }|                     dddd¦  «        }t          j
                             ||	|
fdd	¬
¦  «        }|                     dddd¦  «                             dd|¦  «        }t          j        ||fd¬¦  «        S )a   
        This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher resolution
        images. This method is also adapted to support torch.jit tracing.

        Adapted from:
        - https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174-L194, and
        - https://github.com/facebookresearch/dinov2/blob/e1277af2ba9496fbadf7aec6eba56e8d882d1e35/dinov2/models/vision_transformer.py#L179-L211
        r   Néÿÿÿÿg      à?r   r   é   ÚbicubicF)ÚsizeÚmodeÚalign_corners©Údim)Úshaperr   r9   ÚjitÚ
is_tracingre   r   ÚreshapeÚpermuter   Ú
functionalÚinterpolateÚviewÚcat)r-   rw   rx   ry   rq   Únum_positionsÚclass_pos_embedÚpatch_pos_embedr‚   Ú
new_heightÚ	new_widthÚsqrt_num_positionss               r.   Úinterpolate_pos_encodingz-FlavaImageEmbeddings.interpolate_pos_encoding  sr  € ð !Ô& qÔ)¨AÑ-ˆØÔ0Ô6°qÔ9¸AÑ=ˆõ Œy×#Ò#Ñ%Ô%ð 	,¨+¸Ò*FÐ*FÈ6ÐUZÊ?È?ØÔ+Ð+àÔ2°1°1°1°b°q°b°5Ô9ˆØÔ2°1°1°1°a°b°b°5Ô9ˆàÔ˜rÔ"ˆà˜tœÑ.ˆ
Ø˜Tœ_Ñ,ˆ	å& }°cÑ'9Ñ:Ô:ÐØ)×1Ò1°!Ð5GÐI[Ð]`ÑaÔaˆØ)×1Ò1°!°Q¸¸1Ñ=Ô=ˆåœ-×3Ò3ØØ˜iÐ(ØØð	 4ñ 
ô 
ˆð *×1Ò1°!°Q¸¸1Ñ=Ô=×BÒBÀ1ÀbÈ#ÑNÔNˆåŒy˜/¨?Ð;ÀÐCÑCÔCÐCr0   Úpixel_valuesÚbool_masked_posr’   c                 ó†  — |j         \  }}}}|                      ||¬¦  «        }|                     ¦   «         \  }}	}
|�“| j                             ||	d¦  «        }|                     ¦   «         dk    r)|                     |                     d¦  «        d¦  «        }|                     d¦  «                             |¦  «        }|d|z
  z  ||z  z   }| j	                             |dd¦  «        }t          j        ||fd¬¦  «        }|r||                      |||¦  «        z   }n
|| j        z   }|                      |¦  «        }|S )N)r’   r{   r   r   g      ð?r   r�   )rƒ   rp   r~   rj   Úexpandr‚   rŠ   Ú	unsqueezeÚtype_asrn   r9   r‹   r’   rr   ru   )r-   r“   r”   r’   Ú
batch_sizerf   rx   ry   rw   Úseq_lenÚ_Úmask_tokensÚmaskÚ
cls_tokenss                 r.   ÚforwardzFlavaImageEmbeddings.forward*  sY  € ð 3?Ô2DÑ/ˆ
�L &¨%Ø×*Ò*¨<ÐRjÐ*ÑkÔkˆ
à!+§¢Ñ!2Ô!2Ñˆ
�G˜QØÐ&Øœ/×0Ò0°¸WÀbÑIÔIˆKà×"Ò"Ñ$Ô$¨Ò)Ð)Ø"1×"6Ò"6°×7KÒ7KÈAÑ7NÔ7NÐPRÑ"SÔ"S�à"×,Ò,¨RÑ0Ô0×8Ò8¸ÑEÔEˆDØ# s¨T¡zÑ2°[À4Ñ5GÑGˆJð ”^×*Ò*¨:°r¸2Ñ>Ô>ˆ
Ý”Y 
¨JÐ7¸QÐ?Ñ?Ô?ˆ
ð $ð 	?Ø# d×&CÒ&CÀJÐPVÐX]Ñ&^Ô&^Ñ^ˆJˆJà# dÔ&>Ñ>ˆJà—\’\ *Ñ-Ô-ˆ
àÐr0   ©F©NF)r5   r6   r7   r8   r   rI   ri   r9   ÚTensorÚintr’   Ú
BoolTensorrŸ   Ú__classcell__©rv   s   @r.   r`   r`   ê   sö   ø€ € € € € ðð ðð Ð/ð Àð ÐRVð ð ð ð ð ð ð&&D°5´<ð &DÈð &DÐUXð &DÐ]bÔ]ið &Dð &Dð &Dð &DðV 48Ø).ð	ð à”lðð Ô)¨DÑ0ðð #'ð	ð
 
Œðð ð ð ð ð ð ð r0   r`   c            	       ó¦   ‡ — e Zd ZdZ	 	 	 	 ddeee         z  eeef         z  deeeef         z  ded	efˆ fd
„Zddej	        de
dej	        fd„Zˆ xZS )ro   z#
    Image to Patch Embedding.
    éà   é   r   é   rd   re   rf   rg   c                 ó~  •— t          ¦   «                              ¦   «          t          |t          j        j        ¦  «        s||f}t          |t          j        j        ¦  «        s||f}|d         |d         z  |d         |d         z  z  }|| _        || _        || _        t          j
        ||||¬¦  «        | _        d S )Nr   r   )Úkernel_sizeÚstride)rh   ri   Ú
isinstanceÚcollectionsÚabcÚIterablerd   re   rq   r   ÚConv2dÚ
projection)r-   rd   re   rf   rg   rq   rv   s         €r.   ri   zPatchEmbeddings.__init__S  s·   ø€ õ 	‰Œ×ÒÑÔÐÝ˜*¥k¤oÔ&>Ñ?Ô?ð 	2Ø$ jÐ1ˆJÝ˜*¥k¤oÔ&>Ñ?Ô?ð 	2Ø$ jÐ1ˆJØ! !”}¨
°1¬Ñ5¸*ÀQ¼-È:ÐVWÌ=Ñ:XÑYˆØ$ˆŒØ$ˆŒØ&ˆÔåœ) L°)ÈÐ\fÐgÑgÔgˆŒˆˆr0   Fr“   r’   r%   c                 óB  — |j         \  }}}}|sT|| j        d         k    s|| j        d         k    r2t          d|› d|› d| j        d         › d| j        d         › d�	¦  «        ‚|                      |¦  «                             d¦  «                             dd¦  «        }|S )Nr   r   zInput image size (Ú*z) doesn't match model (z).r|   )rƒ   rd   Ú
ValueErrorr³   ÚflattenÚ	transpose)r-   r“   r’   r™   rf   rx   ry   Úxs           r.   rŸ   zPatchEmbeddings.forwardf  sÚ   € Ø2>Ô2DÑ/ˆ
�L &¨%Ø'ð 	Ø˜œ¨Ô+Ò+Ð+¨u¸¼ÈÔ8JÒ/JÐ/JÝ ðE¨ð Eð E°%ð Eð EØœ¨Ô+ðEð EØ.2¬o¸aÔ.@ðEð Eð Eñô ð ð �OŠO˜LÑ)Ô)×1Ò1°!Ñ4Ô4×>Ò>¸qÀ!ÑDÔDˆØˆr0   )r¨   r©   r   rª   r    )r5   r6   r7   r8   r£   Úlistr2   ri   r9   r¢   rI   rŸ   r¥   r¦   s   @r.   ro   ro   N  sá   ø€ € € € € ðð ð 9<Ø,.ØØðhð hà˜$˜sœ)‘O e¨C°¨H¤oÑ5ðhð ˜%  S œ/Ñ)ðhð ð	hð
 ðhð hð hð hð hð hð&	ð 	 E¤Lð 	ÈDð 	Ð]bÔ]ið 	ð 	ð 	ð 	ð 	ð 	ð 	ð 	r0   ro   c                   ón   ‡ — e Zd ZdZˆ fd„Z	 	 	 ddej        dz  dej        dz  dej        dz  fd„Zˆ xZS )	ÚFlavaTextEmbeddingszGConstruct the embeddings from word, position and token_type embeddings.c                 óÒ  •— t          ¦   «                              ¦   «          t          j        |j        |j        |j        ¬¦  «        | _        t          j        |j        |j        ¦  «        | _	        t          j        |j
        |j        ¦  «        | _        t          j        |j        |j        ¬¦  «        | _        t          j        |j        ¦  «        | _        |                      dt%          j        |j        ¦  «                             d¦  «        d¬¦  «         |                      dt%          j        | j                             ¦   «         t$          j        ¬¦  «        d¬¦  «         d S )	N)Úpadding_idx©ÚepsÚposition_ids©r   r{   F)Ú
persistentÚtoken_type_ids)Údtype)rh   ri   r   Ú	EmbeddingÚ
vocab_sizerm   Úpad_token_idÚword_embeddingsÚmax_position_embeddingsrr   Útype_vocab_sizeÚtoken_type_embeddingsÚ	LayerNormÚlayer_norm_epsrs   rt   ru   Úregister_bufferr9   Úaranger–   rl   rÁ   r~   Úlong©r-   ra   rv   s     €r.   ri   zFlavaTextEmbeddings.__init__u  s3  ø€ Ý‰Œ×ÒÑÔÐÝ!œ|¨FÔ,=¸vÔ?QÐ_eÔ_rÐsÑsÔsˆÔÝ#%¤<°Ô0NÐPVÔPbÑ#cÔ#cˆÔ Ý%'¤\°&Ô2HÈ&ÔJ\Ñ%]Ô%]ˆÔ"åœ fÔ&8¸fÔ>SÐTÑTÔTˆŒÝ”z &Ô"<Ñ=Ô=ˆŒà×ÒØ�EœL¨Ô)GÑHÔH×OÒOÐPWÑXÔXÐejð 	ñ 	
ô 	
ð 	
ð 	×ÒØ�eœk¨$Ô*;×*@Ò*@Ñ*BÔ*BÍ%Ì*ÐUÑUÔUÐbgð 	ñ 	
ô 	
ð 	
ð 	
ð 	
r0   NÚ	input_idsrÄ   rÁ   c                 ó,  — |                      ¦   «         }|d         }|€| j        d d …d |…f         }|€mt          | d¦  «        r2| j        d d …d |…f         }|                     |d         |¦  «        }|}n+t          j        |t
          j        | j        j        ¬¦  «        }|  	                    |¦  «        }|  
                    |¦  «        }	||	z   }
|                      |¦  «        }|
|z  }
|                      |
¦  «        }
|                      |
¦  «        }
|
S )Nr   rÄ   r   )rÅ   Údevice)r~   rÁ   ÚhasattrrÄ   r–   r9   rl   rÑ   rÕ   rÉ   rÌ   rr   rÍ   ru   )r-   rÓ   rÄ   rÁ   Úinput_shapeÚ
seq_lengthÚbuffered_token_type_idsÚ buffered_token_type_ids_expandedÚinputs_embedsrÌ   rw   rr   s               r.   rŸ   zFlavaTextEmbeddings.forward…  s.  € ð  —n’nÑ&Ô&ˆØ  ”^ˆ
àÐØÔ,¨Q¨Q¨Q°°°¨^Ô<ˆLð
 Ð!Ý�tÐ-Ñ.Ô.ð mØ*.Ô*=¸a¸a¸aÀÀ*À¸nÔ*MÐ'Ø3J×3QÒ3QÐR]Ð^_ÔR`ÐblÑ3mÔ3mÐ0Ø!A��å!&¤¨[ÅÄ
ÐSWÔSdÔSkÐ!lÑ!lÔ!l�à×,Ò,¨YÑ7Ô7ˆØ $× :Ò :¸>Ñ JÔ JÐØ"Ð%:Ñ:ˆ
à"×6Ò6°|ÑDÔDÐØÐ)Ñ)ˆ
à—^’^ JÑ/Ô/ˆ
Ø—\’\ *Ñ-Ô-ˆ
ØÐr0   ©NNN)	r5   r6   r7   r8   ri   r9   r¢   rŸ   r¥   r¦   s   @r.   r¼   r¼   r  s“   ø€ € € € € ØQÐQð
ð 
ð 
ð 
ð 
ð$ *.Ø.2Ø,0ð	 ð  à”< $Ñ&ð ð œ tÑ+ð ð ”l TÑ)ð	 ð  ð  ð  ð  ð  ð  ð  r0   r¼   c                   ó    ‡ — e Zd Zdeddfˆ fd„Z	 	 d
dej        dej        dz  dedeej        ej        f         eej                 z  fd	„Z	ˆ xZ
S )ÚFlavaSelfAttentionra   r%   Nc                 óŽ  •— t          ¦   «                              ¦   «          |j        |j        z  dk    r0t	          |d¦  «        s t          d|j        › d|j        › d�¦  «        ‚|j        | _        t          |j        |j        z  ¦  «        | _        | j        | j        z  | _        t          j
        |j        | j        |j        ¬¦  «        | _        t          j
        |j        | j        |j        ¬¦  «        | _        t          j
        |j        | j        |j        ¬¦  «        | _        t          j        |j        ¦  «        | _        d S )Nr   Úembedding_sizezThe hidden size z4 is not a multiple of the number of attention heads ú.©Úbias)rh   ri   rm   Únum_attention_headsrÖ   r¶   r£   Úattention_head_sizeÚall_head_sizer   ÚLinearÚqkv_biasÚqueryÚkeyÚvaluers   Úattention_probs_dropout_probru   rÒ   s     €r.   ri   zFlavaSelfAttention.__init__©  s.  ø€ Ý‰Œ×ÒÑÔÐØÔ Ô :Ñ:¸aÒ?Ð?ÍÐPVÐXhÑHiÔHiÐ?Ýð7 6Ô#5ð 7ð 7ØÔ3ð7ð 7ð 7ñô ð ð
 $*Ô#=ˆÔ Ý#& vÔ'9¸FÔ<VÑ'VÑ#WÔ#WˆÔ Ø!Ô5¸Ô8PÑPˆÔå”Y˜vÔ1°4Ô3EÈFÌOÐ\Ñ\Ô\ˆŒ
Ý”9˜VÔ/°Ô1CÈ&Ì/ÐZÑZÔZˆŒÝ”Y˜vÔ1°4Ô3EÈFÌOÐ\Ñ\Ô\ˆŒ
å”z &Ô"EÑFÔFˆŒˆˆr0   FÚhidden_statesÚattention_maskÚoutput_attentionsc                 óš  — |j         d d…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }t          j        ||                     dd¦  «        ¦  «        }	|	t          j
        | j        ¦  «        z  }	|�|	|z   }	t          j                             |	d¬¦  «        }
|                      |
¦  «        }
t          j        |
|¦  «        }|                     dddd¦  «                             ¦   «         }|                     ¦   «         d d…         | j        fz   } |j        |Ž }|r||
fn|f}|S )Nr{   r   r|   éþÿÿÿr�   r   r   )rƒ   rå   ré   rŠ   r¸   rê   rë   r9   ÚmatmulÚmathÚsqrtr   rˆ   Úsoftmaxru   r‡   Ú
contiguousr~   ræ   )r-   rí   rî   rï   r×   Úhidden_shapeÚquery_layerÚ	key_layerÚvalue_layerÚattention_scoresÚattention_probsÚcontext_layerÚnew_context_layer_shapeÚoutputss                 r.   rŸ   zFlavaSelfAttention.forward»  sÃ  € ð $Ô)¨#¨2¨#Ô.ˆØC˜ÐC bÐC¨$Ô*BÐCÐCˆØ—j’j Ñ/Ô/×4Ò4°\ÑBÔB×LÒLÈQÐPQÑRÔRˆØ—H’H˜]Ñ+Ô+×0Ò0°Ñ>Ô>×HÒHÈÈAÑNÔNˆ	Ø—j’j Ñ/Ô/×4Ò4°\ÑBÔB×LÒLÈQÐPQÑRÔRˆõ !œ<¨°Y×5HÒ5HÈÈRÑ5PÔ5PÑQÔQÐà+­d¬i¸Ô8PÑ.QÔ.QÑQÐØÐ%à/°.Ñ@Ðõ œ-×/Ò/Ð0@ÀbÐ/ÑIÔIˆð Ÿ,š, Ñ7Ô7ˆåœ _°kÑBÔBˆà%×-Ò-¨a°°A°qÑ9Ô9×DÒDÑFÔFˆØ"/×"4Ò"4Ñ"6Ô"6°s¸°sÔ";¸tÔ?QÐ>SÑ"SÐØ*˜Ô*Ð,CÐDˆà6GÐ]�= /Ð2Ð2ÈmÐM]ˆàˆr0   r¡   ©r5   r6   r7   ÚFlavaPossibleConfigsri   r9   r¢   rI   r2   rŸ   r¥   r¦   s   @r.   rÞ   rÞ   ¨  s¾   ø€ € € € € ðGÐ3ð G¸ð Gð Gð Gð Gð Gð Gð* /3Ø"'ð	#ð #à”|ð#ð œ tÑ+ð#ð  ð	#ð
 
ˆuŒ|˜Uœ\Ð)Ô	*¨U°5´<Ô-@Ñ	@ð#ð #ð #ð #ð #ð #ð #ð #r0   rÞ   c                   ó^   ‡ — e Zd ZdZdeddfˆ fd„Zdej        dej        dej        fd„Zˆ xZ	S )	ÚFlavaSelfOutputzµ
    The residual connection is defined in FlavaLayer (same as ViTLayer) instead of here (as is the case with other
    models), due to the layernorm applied before each block.
    ra   r%   Nc                 óÌ   •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          j        |j        ¦  «        | _        d S r]   )	rh   ri   r   rç   rm   Údensers   rt   ru   rÒ   s     €r.   ri   zFlavaSelfOutput.__init__ç  sJ   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3EÑFÔFˆŒ
Ý”z &Ô"<Ñ=Ô=ˆŒˆˆr0   rí   Úinput_tensorc                 óZ   — |                       |¦  «        }|                      |¦  «        }|S r]   ©r  ru   ©r-   rí   r  s      r.   rŸ   zFlavaSelfOutput.forwardì  s*   € ØŸ
š
 =Ñ1Ô1ˆØŸš ]Ñ3Ô3ˆàÐr0   )
r5   r6   r7   r8   r  ri   r9   r¢   rŸ   r¥   r¦   s   @r.   r  r  á  s‡   ø€ € € € € ðð ð
>Ð3ð >¸ð >ð >ð >ð >ð >ð >ð
 U¤\ð ÀÄð ÐRWÔR^ð ð ð ð ð ð ð ð r0   r  c                   ó    ‡ — e Zd Zdeddfˆ fd„Z	 	 d
dej        dej        dz  dedeej        ej        f         eej                 z  fd	„Z	ˆ xZ
S )ÚFlavaAttentionra   r%   Nc                 ó˜   •— t          ¦   «                              ¦   «          t          |¦  «        | _        t	          |¦  «        | _        d S r]   )rh   ri   rÞ   Ú	attentionr  ÚoutputrÒ   s     €r.   ri   zFlavaAttention.__init__ô  s;   ø€ Ý‰Œ×ÒÑÔÐÝ+¨FÑ3Ô3ˆŒÝ% fÑ-Ô-ˆŒˆˆr0   Frí   rî   rï   c                 óŠ   — |                       |||¬¦  «        }|                      |d         |¦  «        }|f|dd …         z   }|S ©N)rî   rï   r   r   )r  r  )r-   rí   rî   rï   Úself_outputsÚattention_outputrÿ   s          r.   rŸ   zFlavaAttention.forwardù  sY   € ð —~’~Ø¨.ÐL]ð &ñ 
ô 
ˆð  Ÿ;š; |°A¤¸ÑFÔFÐà#Ð%¨°Q°R°RÔ(8Ñ8ˆØˆr0   r¡   r   r¦   s   @r.   r  r  ó  s¶   ø€ € € € € ð.Ð3ð .¸ð .ð .ð .ð .ð .ð .ð /3Ø"'ð	ð à”|ðð œ tÑ+ðð  ð	ð
 
ˆuŒ|˜Uœ\Ð)Ô	*¨U°5´<Ô-@Ñ	@ðð ð ð ð ð ð ð r0   r  c                   óL   ‡ — e Zd Zdeddfˆ fd„Zdej        dej        fd„Zˆ xZS )ÚFlavaIntermediatera   r%   Nc                 ó  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          |j        t          ¦  «        rt          |j                 | _        d S |j        | _        d S r]   )rh   ri   r   rç   rm   Úintermediate_sizer  r®   Ú
hidden_actÚstrr	   Úintermediate_act_fnrÒ   s     €r.   ri   zFlavaIntermediate.__init__
  sn   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3KÑLÔLˆŒ
Ý�fÔ'­Ñ-Ô-ð 	9Ý'-¨fÔ.?Ô'@ˆDÔ$Ð$Ð$à'-Ô'8ˆDÔ$Ð$Ð$r0   rí   c                 óZ   — |                       |¦  «        }|                      |¦  «        }|S r]   )r  r  ©r-   rí   s     r.   rŸ   zFlavaIntermediate.forward  s,   € ØŸ
š
 =Ñ1Ô1ˆØ×0Ò0°Ñ?Ô?ˆàÐr0   ©	r5   r6   r7   r  ri   r9   r¢   rŸ   r¥   r¦   s   @r.   r  r  	  sr   ø€ € € € € ð9Ð3ð 9¸ð 9ð 9ð 9ð 9ð 9ð 9ð U¤\ð °e´lð ð ð ð ð ð ð ð r0   r  c                   óZ   ‡ — e Zd Zdeddfˆ fd„Zdej        dej        dej        fd„Zˆ xZS )ÚFlavaOutputra   r%   Nc                 óÌ   •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          j        |j        ¦  «        | _	        d S r]   )
rh   ri   r   rç   r  rm   r  rs   rt   ru   rÒ   s     €r.   ri   zFlavaOutput.__init__  sJ   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ7¸Ô9KÑLÔLˆŒ
Ý”z &Ô"<Ñ=Ô=ˆŒˆˆr0   rí   r  c                 ód   — |                       |¦  «        }|                      |¦  «        }||z   }|S r]   r  r	  s      r.   rŸ   zFlavaOutput.forward!  s4   € ØŸ
š
 =Ñ1Ô1ˆØŸš ]Ñ3Ô3ˆà%¨Ñ4ˆàÐr0   r  r¦   s   @r.   r  r    s}   ø€ € € € € ð>Ð3ð >¸ð >ð >ð >ð >ð >ð >ð U¤\ð ÀÄð ÐRWÔR^ð ð ð ð ð ð ð ð r0   r  c                   ó¤   ‡ — e Zd ZdZdeddfˆ fd„Z	 	 ddej        dej        dz  d	ede	ej        ej        f         e	ej                 z  fd
„Z
ˆ xZS )Ú
FlavaLayerz?This corresponds to the Block class in the timm implementation.ra   r%   Nc                 óz  •— t          ¦   «                              ¦   «          |j        | _        d| _        t	          |¦  «        | _        t          |¦  «        | _        t          |¦  «        | _	        t          j        |j        |j        ¬¦  «        | _        t          j        |j        |j        ¬¦  «        | _        d S )Nr   r¿   )rh   ri   Úchunk_size_feed_forwardÚseq_len_dimr  r  r  Úintermediater  r  r   rÍ   rm   rÎ   Úlayernorm_beforeÚlayernorm_afterrÒ   s     €r.   ri   zFlavaLayer.__init__-  sœ   ø€ Ý‰Œ×ÒÑÔÐØ'-Ô'EˆÔ$ØˆÔÝ'¨Ñ/Ô/ˆŒÝ-¨fÑ5Ô5ˆÔÝ! &Ñ)Ô)ˆŒõ !#¤¨VÔ-?ÀVÔEZÐ [Ñ [Ô [ˆÔÝ!œ|¨FÔ,>ÀFÔDYÐZÑZÔZˆÔÐÐr0   Frí   rî   rï   c                 ó  — |                       |                      |¦  «        ||¬¦  «        }|d         }|dd …         }||z   }|                      |¦  «        }|                      |¦  «        }|                      ||¦  «        }|f|z   }|S r  )r  r'  r(  r&  r  )r-   rí   rî   rï   Úself_attention_outputsr  rÿ   Úlayer_outputs           r.   rŸ   zFlavaLayer.forward9  s©   € ð "&§¢Ø×!Ò! -Ñ0Ô0Ø)Ø/ð "0ñ "
ô "
Ðð
 2°!Ô4ÐØ(¨¨¨Ô,ˆð )¨=Ñ8ˆð ×+Ò+¨MÑ:Ô:ˆØ×(Ò(¨Ñ6Ô6ˆð —{’{ <°Ñ?Ô?ˆà�/ GÑ+ˆàˆr0   r¡   )r5   r6   r7   r8   r  ri   r9   r¢   rI   r2   rŸ   r¥   r¦   s   @r.   r"  r"  *  sÄ   ø€ € € € € ØIÐIð
[Ð3ð 
[¸ð 
[ð 
[ð 
[ð 
[ð 
[ð 
[ð /3Ø"'ð	ð à”|ðð œ tÑ+ðð  ð	ð
 
ˆuŒ|˜Uœ\Ð)Ô	*¨U°5´<Ô-@Ñ	@ðð ð ð ð ð ð ð r0   r"  c                   ór   ‡ — e Zd Zdeddfˆ fd„Z	 	 	 	 ddej        dej        dz  d	ed
ededee	z  fd„Z
ˆ xZS )ÚFlavaEncoderra   r%   Nc                 óÔ   •‡— t          ¦   «                              ¦   «          ‰| _        t          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        d| _        d S )Nc                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r<   )r"  )r+   r›   ra   s     €r.   ú
<listcomp>z)FlavaEncoder.__init__.<locals>.<listcomp>Z  s!   ø€ Ð#`Ð#`Ð#`¸1¥J¨vÑ$6Ô$6Ð#`Ð#`Ð#`r0   F)	rh   ri   ra   r   Ú
ModuleListÚrangeÚnum_hidden_layersÚlayerÚgradient_checkpointingrÒ   s    `€r.   ri   zFlavaEncoder.__init__W  s`   øø€ Ý‰Œ×ÒÑÔÐØˆŒÝ”]Ð#`Ð#`Ð#`Ð#`ÅÀfÔF^Ñ@_Ô@_Ð#`Ñ#`Ô#`ÑaÔaˆŒ
Ø&+ˆÔ#Ð#Ð#r0   FTrí   rî   rï   Úoutput_hidden_statesÚreturn_dictc                 ó  — |rdnd }|rdnd }t          | j        ¦  «        D ]0\  }}	|r||fz   } |	|||¦  «        }
|
d         }|r||
d         fz   }Œ1|r||fz   }|st          d„ |||fD ¦   «         ¦  «        S t          |||¬¦  «        S )Nr<   r   r   c              3   ó   K  — | ]}|®|V — Œ	d S r]   r<   )r+   rH   s     r.   r/   z'FlavaEncoder.forward.<locals>.<genexpr>w  s(   è è € ÐmÐm˜qÐ_`Ð_l˜Ð_lÐ_lÐ_lÐ_lÐmÐmr0   )Úlast_hidden_staterí   Ú
attentions)Ú	enumerater4  r2   r   )r-   rí   rî   rï   r6  r7  Úall_hidden_statesÚall_self_attentionsÚiÚlayer_moduleÚlayer_outputss              r.   rŸ   zFlavaEncoder.forward]  sù   € ð #7Ð@˜B˜B¸DÐØ$5Ð?˜b˜b¸4Ðå(¨¬Ñ4Ô4ð 		Pð 		P‰OˆAˆ|Ø#ð IØ$5¸Ð8HÑ$HÐ!à(˜L¨¸ÐHYÑZÔZˆMà)¨!Ô,ˆMà ð PØ&9¸]È1Ô=MÐ<OÑ&OÐ#øàð 	EØ 1°]Ð4DÑ DÐàð 	nÝÐmÐm ]Ð4EÐGZÐ$[ÐmÑmÔmÑmÔmÐmÝØ+Ð;LÐYlð
ñ 
ô 
ð 	
r0   )NFFT)r5   r6   r7   r   ri   r9   r¢   rI   r2   r   rŸ   r¥   r¦   s   @r.   r-  r-  V  sº   ø€ € € € € ð,˜{ð ,¨tð ,ð ,ð ,ð ,ð ,ð ,ð /3Ø"'Ø%*Ø ð
ð 
à”|ð
ð œ tÑ+ð
ð  ð	
ð
 #ð
ð ð
ð 
�Ñ	 ð
ð 
ð 
ð 
ð 
ð 
ð 
ð 
r0   r-  c                   ó:   ‡ — e Zd Zdefˆ fd„Zdej        fd„Zˆ xZS )ÚFlavaPoolerra   c                 óÀ   •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          j        ¦   «         | _        d S r]   )rh   ri   r   rç   rm   r  ÚTanhÚ
activationrÒ   s     €r.   ri   zFlavaPooler.__init__~  sC   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3EÑFÔFˆŒ
Ýœ'™)œ)ˆŒˆˆr0   rí   c                 ór   — |d d …df         }|                       |¦  «        }|                      |¦  «        }|S ©Nr   )r  rF  )r-   rí   Úfirst_token_tensorÚpooled_outputs       r.   rŸ   zFlavaPooler.forwardƒ  s@   € ð +¨1¨1¨1¨a¨4Ô0ÐØŸ
š
Ð#5Ñ6Ô6ˆØŸš¨Ñ6Ô6ˆØÐr0   r  r¦   s   @r.   rC  rC  }  sb   ø€ € € € € ð$Ð3ð $ð $ð $ð $ð $ð $ð
 U¤\ð ð ð ð ð ð ð ð r0   rC  c                   ó”   ‡ — e Zd ZU eed<   dZdZdZ ej	        ¦   «         de
j        e
j        z  e
j        z  ddfˆ fd„¦   «         Zˆ xZS )	ÚFlavaPreTrainedModelra   Úflava)ÚimageÚtextTÚmoduler%   Nc                 óf  •— t          ¦   «                              |¦  «         t          |t          ¦  «        rt	          j        |j        ¦  «         dS t          |t          ¦  «        rVt	          j        |j        ¦  «         t	          j        |j	        ¦  «         |j
        �t	          j        |j
        ¦  «         dS dS t          |t          ¦  «        rjt	          j        |j        t          j        |j        j        d         ¦  «                             d¦  «        ¦  «         t	          j        |j        ¦  «         dS t          |t&          ¦  «        r$|j        rt	          j        |j        ¦  «         dS dS t          |t*          ¦  «        r&t	          j        |j        | j        j        ¦  «         dS dS )zInitialize the weightsNr{   rÂ   )rh   Ú_init_weightsr®   ÚFlavaMaskedPredictionHeadÚinitÚzeros_rã   r`   rn   rr   rj   r¼   Úcopy_rÁ   r9   rÐ   rƒ   r–   rÄ   ÚFlavaMultimodalModelÚuse_cls_tokenÚ
FlavaModelÚ	constant_Úlogit_scalera   Úlogit_scale_init_value)r-   rP  rv   s     €r.   rR  z"FlavaPreTrainedModel._init_weights“  s•  ø€ õ 	‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ7Ñ8Ô8ð 	SÝŒK˜œÑ$Ô$Ð$Ð$Ð$Ý˜Õ 4Ñ5Ô5ð 	SÝŒK˜Ô(Ñ)Ô)Ð)ÝŒK˜Ô2Ñ3Ô3Ð3ØÔ Ð,Ý”˜FÔ-Ñ.Ô.Ð.Ð.Ð.ð -Ð,å˜Õ 3Ñ4Ô4ð 	SÝŒJ�vÔ*­E¬L¸Ô9LÔ9RÐSUÔ9VÑ,WÔ,W×,^Ò,^Ð_fÑ,gÔ,gÑhÔhÐhÝŒK˜Ô-Ñ.Ô.Ð.Ð.Ð.Ý˜Õ 4Ñ5Ô5ð 	SØÔ#ð .Ý”˜FÔ,Ñ-Ô-Ð-Ð-Ð-ð.ð .å˜¥
Ñ+Ô+ð 	SÝŒN˜6Ô-¨t¬{Ô/QÑRÔRÐRÐRÐRð	Sð 	Sr0   )r5   r6   r7   r   r;   Úbase_model_prefixÚinput_modalitiesÚsupports_gradient_checkpointingr9   Úno_gradr   rç   r²   rÍ   rR  r¥   r¦   s   @r.   rL  rL  Œ  s™   ø€ € € € € € àÐÐÑØÐØ(ÐØ&*Ð#à€U„]�_„_ðS B¤I°´	Ñ$9¸B¼LÑ$Hð SÈTð Sð Sð Sð Sð Sñ „_ðSð Sð Sð Sð Sr0   rL  c                   ó  ‡ — e Zd ZU eed<   dZdZdZddedefˆ fd„Z	de
j        fd	„Zd
e
j        fd„Ze	 	 	 	 	 	 	 ddej        dz  dej        dz  dedz  dej        dz  dedz  dedz  dedz  deez  fd„¦   «         Zˆ xZS )ÚFlavaImageModelra   zflava.image_modelr“   ©rN  TÚadd_pooling_layerc                 óJ  •— t          ¦   «                              |¦  «         || _        t          |¦  «        | _        t          |¦  «        | _        t          j        |j	        |j
        ¬¦  «        | _        |rt          |¦  «        nd| _        |                      ¦   «          dS ©úv
        add_pooling_layer (bool, *optional*, defaults to `True`):
            Whether to add a pooling layer
        r¿   N)rh   ri   ra   r`   rw   r-  Úencoderr   rÍ   rm   rÎ   Ú	layernormrC  ÚpoolerÚ	post_init©r-   ra   rd  rv   s      €r.   ri   zFlavaImageModel.__init__°  s�   ø€ õ
 	‰Œ×Ò˜Ñ Ô Ð àˆŒå.¨vÑ6Ô6ˆŒÝ# FÑ+Ô+ˆŒåœ fÔ&8¸fÔ>SÐTÑTÔTˆŒØ->ÐH•k &Ñ)Ô)Ð)ÀDˆŒà�ŠÑÔÐÐÐr0   r%   c                 ó   — | j         j        S r]   ©rw   rp   r4   s    r.   Úget_input_embeddingsz$FlavaImageModel.get_input_embeddingsÁ  s   € ØŒÔ/Ð/r0   rë   c                 ó   — || j         _        d S r]   rn  ©r-   rë   s     r.   Úset_input_embeddingsz$FlavaImageModel.set_input_embeddingsÄ  s   € Ø+0ˆŒÔ(Ð(Ð(r0   Nr”   r’   rî   rï   r6  r7  c                 óº  — |�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|€t	          d¦  «        ‚|                      |||¬¦  «        }	|                      |	||||¬¦  «        }
|
d         }|                      |¦  «        }| j        �|                      |¦  «        nd}|s||f|
dd…         z   S t          |||
j
        |
j        ¬¦  «        S )zÅ
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, image_num_patches)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).
        Nz You have to specify pixel_values)r”   r’   ©rî   rï   r6  r7  r   r   ©r:  Úpooler_outputrí   r;  )ra   rï   r6  r7  r¶   rw   rh  ri  rj  r   rí   r;  )r-   r“   r”   r’   rî   rï   r6  r7  ÚkwargsÚembedding_outputÚencoder_outputsÚsequence_outputrJ  s                r.   rŸ   zFlavaImageModel.forwardÇ  s/  € ð  2CÐ1NÐ-Ð-ÐTXÔT_ÔTqÐà$8Ð$DÐ Ð È$Ì+ÔJjð 	ð &1Ð%<�k�kÀ$Ä+ÔBYˆàÐÝÐ?Ñ@Ô@Ð@àŸ?š?Ø¨/ÐTlð +ñ 
ô 
Ðð Ÿ,š,ØØ)Ø/Ø!5Ø#ð 'ñ 
ô 
ˆð *¨!Ô,ˆØŸ.š.¨Ñ9Ô9ˆØ8<¼Ð8O˜Ÿš OÑ4Ô4Ð4ÐUYˆàð 	JØ# ]Ð3°oÀaÀbÀbÔ6IÑIÐIå)Ø-Ø'Ø)Ô7Ø&Ô1ð	
ñ 
ô 
ð 	
r0   ©T©NNNNNNN)r5   r6   r7   r   r;   r]  Úmain_input_namer^  rI   ri   r   ÚModulero  rr  r   r9   r¢   r¤   r2   r   rŸ   r¥   r¦   s   @r.   rb  rb  ¨  s`  ø€ € € € € € àÐÐÑà+ÐØ$€OØ!Ððð Ð/ð ÀDð ð ð ð ð ð ð"0 b¤ið 0ð 0ð 0ð 0ð1¨"¬)ð 1ð 1ð 1ð 1ð ð -1Ø37Ø04Ø.2Ø)-Ø,0Ø#'ð/
ð /
à”l TÑ)ð/
ð Ô)¨DÑ0ð/
ð #'¨¡+ð	/
ð
 œ tÑ+ð/
ð   $™;ð/
ð # T™kð/
ð ˜D‘[ð/
ð 
Ð+Ñ	+ð/
ð /
ð /
ñ „^ð/
ð /
ð /
ð /
ð /
r0   rb  c                   ó   ‡ — e Zd ZU eed<   dZdZddedefˆ fd„Zde	fd„Z
d	ej        fd
„Ze	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dedz  dedz  dedz  deez  fd„¦   «         Zˆ xZS )ÚFlavaTextModelra   zflava.text_model)rO  Trd  c                 óJ  •— t          ¦   «                              |¦  «         || _        t          |¦  «        | _        t          |¦  «        | _        t          j        |j	        |j
        ¬¦  «        | _        |rt          |¦  «        nd| _        |                      ¦   «          dS rf  )rh   ri   ra   r¼   rw   r-  rh  r   rÍ   rm   rÎ   ri  rC  rj  rk  rl  s      €r.   ri   zFlavaTextModel.__init__  s�   ø€ õ
 	‰Œ×Ò˜Ñ Ô Ð ØˆŒå-¨fÑ5Ô5ˆŒÝ# FÑ+Ô+ˆŒåœ fÔ&8¸fÔ>SÐTÑTÔTˆŒØ->ÐH•k &Ñ)Ô)Ð)ÀDˆŒà�ŠÑÔÐÐÐr0   r%   c                 ó   — | j         j        S r]   ©rw   rÉ   r4   s    r.   ro  z#FlavaTextModel.get_input_embeddings  s   € ØŒÔ.Ð.r0   rë   c                 ó   — || j         _        d S r]   rƒ  rq  s     r.   rr  z#FlavaTextModel.set_input_embeddings  s   € Ø*/ˆŒÔ'Ð'Ð'r0   NrÓ   rî   rÄ   rÁ   rï   r6  r7  c                 óè  — |�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|€t	          d¦  «        ‚|                      |||¬¦  «        }	t          | j         |	|¬¦  «        }|                      |	||||¬¦  «        }
|
d         }|                      |¦  «        }| j	        �|  	                    |¦  «        nd}|s||f|
dd…         z   S t          |||
j        |
j        ¬¦  «        S )	aù  
        input_ids (`torch.LongTensor` of shape `(batch_size, text_seq_length)`):
            Indices of input sequence tokens in the vocabulary. Indices can be obtained using [`AutoTokenizer`]. See
            [`PreTrainedTokenizer.encode`] and [`PreTrainedTokenizer.__call__`] for details. [What are input
            IDs?](../glossary#input-ids)
        token_type_ids (`torch.LongTensor` of shape `(batch_size, text_seq_length)`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:
            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.
            [What are token type IDs?](../glossary#token-type-ids)
        NzYou have to specify input_ids)rÓ   rÄ   rÁ   ©ra   rÛ   rî   rt  r   r   ru  )ra   rï   r6  r7  r¶   rw   r
   rh  ri  rj  r   rí   r;  )r-   rÓ   rî   rÄ   rÁ   rï   r6  r7  rw  rx  ry  rz  rJ  s                r.   rŸ   zFlavaTextModel.forward  sQ  € ð0 2CÐ1NÐ-Ð-ÐTXÔT_ÔTqÐà$8Ð$DÐ Ð È$Ì+ÔJjð 	ð &1Ð%<�k�kÀ$Ä+ÔBYˆàÐÝÐ<Ñ=Ô=Ð=àŸ?š?ØØ)Ø%ð +ñ 
ô 
Ðõ 3Ø”;Ø*Ø)ð
ñ 
ô 
ˆð Ÿ,š,ØØ)Ø/Ø!5Ø#ð 'ñ 
ô 
ˆð *¨!Ô,ˆØŸ.š.¨Ñ9Ô9ˆØ8<¼Ð8O˜Ÿš OÑ4Ô4Ð4ÐUYˆàð 	JØ# ]Ð3°oÀaÀbÀbÔ6IÑIÐIå)Ø-Ø'Ø)Ô7Ø&Ô1ð	
ñ 
ô 
ð 	
r0   r{  r|  )r5   r6   r7   r   r;   r]  r^  rI   ri   ro   ro  r   r~  rr  r   r9   r¢   r2   r   rŸ   r¥   r¦   s   @r.   r€  r€  ú  sZ  ø€ € € € € € àÐÐÑà*ÐØ Ððð ˜ð À4ð ð ð ð ð ð ð / oð /ð /ð /ð /ð0¨"¬)ð 0ð 0ð 0ð 0ð ð *.Ø.2Ø.2Ø,0Ø)-Ø,0Ø#'ð?
ð ?
à”< $Ñ&ð?
ð œ tÑ+ð?
ð œ tÑ+ð	?
ð
 ”l TÑ)ð?
ð   $™;ð?
ð # T™kð?
ð ˜D‘[ð?
ð 
Ð+Ñ	+ð?
ð ?
ð ?
ñ „^ð?
ð ?
ð ?
ð ?
ð ?
r0   r€  c                   ó¦   ‡ — e Zd ZU eed<   dZdZddefˆ fd„Ze	 	 	 	 dde	j
        de	j
        dz  dedz  d	edz  d
edz  deez  fd„¦   «         Zˆ xZS )rW  ra   zflava.multimodal_modelrí   Tc                 ó¶  •— t          ¦   «                              |¦  «         || _        | j        j        | _        | j        r2t	          j        t          j        dd|j        ¦  «        ¦  «        | _	        t          |¦  «        | _        t	          j        |j        |j        ¬¦  «        | _        |rt          |¦  «        nd| _        |                      ¦   «          dS )rg  r   r¿   N)rh   ri   ra   rX  r   rk   r9   rl   rm   rn   r-  rh  rÍ   rÎ   ri  rC  rj  rk  rl  s      €r.   ri   zFlavaMultimodalModel.__init__a  s¹   ø€ õ
 	‰Œ×Ò˜Ñ Ô Ð ØˆŒØ!œ[Ô6ˆÔØÔð 	QÝœ\­%¬+°a¸¸FÔ<NÑ*OÔ*OÑPÔPˆDŒNå# FÑ+Ô+ˆŒåœ fÔ&8¸fÔ>SÐTÑTÔTˆŒØ->ÐH•k &Ñ)Ô)Ð)ÀDˆŒà�ŠÑÔÐÐÐr0   Nrî   rï   r6  r7  r%   c                 óF  — |�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|                     ¦   «         \  }}}	| j        r9| j                             |dd¦  «        }
t          j	        |
|fd¬¦  «        }|dz  }t          | j         ||¬¦  «        }|                      |||||¬¦  «        }|d         }|                      |¦  «        }| j        �|                      |¦  «        nd}|s||f|dd…         z   S t          |||j        |j        ¬¦  «        S )	z¾
        hidden_states (`torch.FloatTensor` of shape `(batch_size, image_num_patches + text_seq_len, hidden_size)`):
            The concatenated hidden states of unimodal encoders.
        Nr{   r   r�   r†  rt  r   ru  )ra   rï   r6  r7  r~   rX  rn   r–   r9   r‹   r
   rh  ri  rj  r   rí   r;  )r-   rí   rî   rï   r6  r7  rw  r™   rØ   r›   rž   ry  rz  rJ  s                 r.   rŸ   zFlavaMultimodalModel.forwards  sv  € ð 2CÐ1NÐ-Ð-ÐTXÔT_ÔTqÐà$8Ð$DÐ Ð È$Ì+ÔJjð 	ð &1Ð%<�k�kÀ$Ä+ÔBYˆà$1×$6Ò$6Ñ$8Ô$8Ñ!ˆ
�J àÔð 	Øœ×.Ò.¨z¸2¸rÑBÔBˆJÝ!œI z°=Ð&AÀqÐIÑIÔIˆMØ˜!‰OˆJå2Ø”;Ø'Ø)ð
ñ 
ô 
ˆð Ÿ,š,ØØ)Ø/Ø!5Ø#ð 'ñ 
ô 
ˆð *¨!Ô,ˆØŸ.š.¨Ñ9Ô9ˆØ8<¼Ð8O˜Ÿš OÑ4Ô4Ð4ÐUYˆàð 	JØ# ]Ð3°oÀaÀbÀbÔ6IÑIÐIå)Ø-Ø'Ø)Ô7Ø&Ô1ð	
ñ 
ô 
ð 	
r0   r{  )NNNN)r5   r6   r7   r   r;   r]  r}  ri   r   r9   r¢   rI   r2   r   rŸ   r¥   r¦   s   @r.   rW  rW  Z  së   ø€ € € € € € à!Ð!Ð!Ñ!à0ÐØ%€Oðð Ð4ð ð ð ð ð ð ð$ ð /3Ø)-Ø,0Ø#'ð3
ð 3
à”|ð3
ð œ tÑ+ð3
ð   $™;ð	3
ð
 # T™kð3
ð ˜D‘[ð3
ð 
Ð+Ñ	+ð3
ð 3
ð 3
ñ „^ð3
ð 3
ð 3
ð 3
ð 3
r0   rW  c                   ó6  ‡ — e Zd ZU eed<   defˆ fd„Zee	 	 	 ddej	        dej	        dz  dej	        dz  dej	        dz  de
e         d	eez  fd
„¦   «         ¦   «         Zee	 	 	 ddej	        dej        dz  dedz  dej	        dz  de
e         d	eez  fd„¦   «         ¦   «         Ze	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej	        dz  dej	        dz  dej	        dz  dej        dz  dej	        dz  dedz  dedz  dededz  d	eez  fd„¦   «         Zˆ xZS )rY  ra   c                 ó~  •— t          ¦   «                              |¦  «         t          |j        t          ¦  «        s%t          dt          |j        ¦  «        › d�¦  «        ‚t          |j        t          ¦  «        s%t          dt          |j        ¦  «        › d�¦  «        ‚t          |j	        t          ¦  «        s(t          ddt          |j	        ¦  «        › d�z   ¦  «        ‚|j        }|j        }|j	        }|j        | _        |j        | _        |j        | _        |j        | _        t!          |¦  «        | _        t%          |¦  «        | _        t)          |¦  «        | _        t-          j        | j        | j        ¦  «        | _        t-          j        | j        | j        ¦  «        | _        t-          j        t7          j        | j        j        ¦  «        ¦  «        | _        t-          j        | j        | j        ¦  «        | _         t-          j        | j        | j        ¦  «        | _!        |  "                    ¦   «          d S )NzLconfig.text_config is expected to be of type FlavaTextConfig but is of type rá   zNconfig.image_config is expected to be of type FlavaImageConfig but is of type zMconfig.multimodal_config is expected to be of type FlavaMultimodalConfig but zis of type )#rh   ri   r®   Útext_configr   Ú	TypeErrorÚtypeÚimage_configr   Úmultimodal_configr   Úprojection_dimrm   Útext_hidden_sizeÚimage_hidden_sizeÚmm_hidden_sizer€  Ú
text_modelrb  Úimage_modelrW  Úmultimodal_modelr   rç   Úimage_projectionÚtext_projectionrk   r9   Útensorra   r\  r[  Úimage_to_mm_projectionÚtext_to_mm_projectionrk  )r-   ra   rŒ  r�  r�  rv   s        €r.   ri   zFlavaModel.__init__®  s  ø€ Ý‰Œ×Ò˜Ñ Ô Ð å˜&Ô,­oÑ>Ô>ð 	Ýð0Ý˜Ô+Ñ,Ô,ð0ð 0ð 0ñô ð õ
 ˜&Ô-Õ/?Ñ@Ô@ð 	Ýð1Ý˜Ô,Ñ-Ô-ð1ð 1ð 1ñô ð õ
 ˜&Ô2Õ4IÑJÔJð 	ÝØ_ØA¥ VÔ%=Ñ >Ô >ÐAÐAÐAñBñô ð ð
 Ô(ˆØÔ*ˆØ"Ô4Ðà$Ô3ˆÔØ +Ô 7ˆÔØ!-Ô!9ˆÔØ/Ô;ˆÔå(¨Ñ5Ô5ˆŒÝ*¨<Ñ8Ô8ˆÔÝ 4Ð5FÑ GÔ GˆÔå "¤	¨$Ô*@À$ÔBUÑ VÔ VˆÔÝ!œy¨Ô)>ÀÔ@SÑTÔTˆÔÝœ<­¬°T´[Ô5WÑ(XÔ(XÑYÔYˆÔå&(¤i°Ô0FÈÔH[Ñ&\Ô&\ˆÔ#Ý%'¤Y¨tÔ/DÀdÔFYÑ%ZÔ%ZˆÔ"à�ŠÑÔÐÐÐr0   NrÓ   rî   rÄ   rÁ   rw  r%   c           	      ón   —  | j         d||||ddœ|¤Ž}|j        }|                      |¦  «        |_        |S )a	  
        input_ids (`torch.LongTensor` of shape `(batch_size, text_seq_length)`):
            Indices of input sequence tokens in the vocabulary. Indices can be obtained using [`AutoTokenizer`]. See
            [`PreTrainedTokenizer.encode`] and [`PreTrainedTokenizer.__call__`] for details. [What are input
            IDs?](../glossary#input-ids)
        token_type_ids (`torch.LongTensor` of shape `(batch_size, text_seq_length)`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:
            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.
            [What are token type IDs?](../glossary#token-type-ids)

        Examples:

        ```python
        >>> import torch
        >>> from transformers import AutoProcessor, FlavaModel

        >>> model = FlavaModel.from_pretrained("{0}")
        >>> processor = AutoProcessor.from_pretrained("{0}")

        >>> inputs = processor(
        ...     text=["a photo of a cat", "a photo of a dog"], max_length=77, padding="max_length", return_tensors="pt"
        ... )
        >>> with torch.inference_mode():
        ...     text_features = model.get_text_features(**inputs)
        ```
        T)rÓ   rî   rÄ   rÁ   r7  r<   )r•  r:  r™  rv  )r-   rÓ   rî   rÄ   rÁ   rw  Útext_outputsr:  s           r.   Úget_text_featureszFlavaModel.get_text_featuresÙ  sd   € ðL 4C°4´?ð 4
ØØ)Ø)Ø%Øð4
ð 4
ð ð4
ð 4
ˆð )Ô:ÐØ%)×%9Ò%9Ð:KÑ%LÔ%LˆÔ"àÐr0   r“   r”   r’   c           	      ón   —  | j         d||||ddœ|¤Ž}|j        }|                      |¦  «        |_        |S )a   
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, image_num_patches)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).

        Examples:

        ```python
        >>> import torch
        >>> from transformers import AutoProcessor, FlavaModel
        >>> from transformers.image_utils import load_image

        >>> model = FlavaModel.from_pretrained("{0}")
        >>> processor = AutoProcessor.from_pretrained("{0}")

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = load_image(url)

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

        >>> with torch.inference_mode():
        ...     image_features = model.get_image_features(**inputs)
        ```
        T)r“   r”   rî   r’   r7  r<   )r–  r:  r˜  rv  )r-   r“   r”   r’   rî   rw  Úimage_outputsr:  s           r.   Úget_image_featureszFlavaModel.get_image_features  se   € ðB 5E°DÔ4Dð 5
Ø%Ø+Ø)Ø%=Øð5
ð 5
ð ð5
ð 5
ˆð *Ô;ÐØ&*×&;Ò&;Ð<MÑ&NÔ&NˆÔ#àÐr0   TÚimage_attention_maskÚskip_multimodal_encoderrï   r6  r7  c           	      óò  — |�|n| j         j        }|
st          d¦  «        ‚d}d}d}d}|�F|                      ||||	|
|¬¦  «        }|d         |d         }}|                      |d         ¦  «        }d}d}d}d}|�G|                      |||||	|
|¬¦  «        }|d         |d         }}|                      |d         ¦  «        }d}d}|�‘|��|s�|�Q|j        \  }}}| j        j	        r|dz  }t          j        |||j        ¬	¦  «        }t          j        ||gd¬
¦  «        }nd}t          j        ||gd¬
¦  «        }|                      |||¬¦  «        }|d         }|s||||||fS t          ||||||¬¦  «        S )a/
  
        input_ids (`torch.LongTensor` of shape `(batch_size, image_num_patches + text_seq_len)`):
            Indices of input sequence tokens in the vocabulary. Indices can be obtained using [`AutoTokenizer`]. See
            [`PreTrainedTokenizer.encode`] and [`PreTrainedTokenizer.__call__`] for details. [What are input
            IDs?](../glossary#input-ids)
        token_type_ids (`torch.LongTensor` of shape `(batch_size, image_num_patches + text_seq_len)`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:
            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.
            [What are token type IDs?](../glossary#token-type-ids)
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, image_num_patches)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).
        image_attention_mask (`torch.Tensor` of shape `(batch_size, image_num_patches)`, *optional*):
            Mask to avoid performing attention on padding pixel values for image inputs. Mask values selected in `[0, 1]`:
            - 1 for pixel values that are real (i.e., **not masked**),
            - 0 for pixel values that are padding (i.e., **masked**).
        skip_multimodal_encoder (*bool*, *optional*):
            Skip any calculations for multimodal encoder. Useful if multimodal encoding is not going to be used.

        Examples:

        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoProcessor, FlavaModel

        >>> model = FlavaModel.from_pretrained("facebook/flava-full")
        >>> processor = AutoProcessor.from_pretrained("facebook/flava-full")

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))

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

        >>> outputs = model(**inputs)

        >>> image_embeddings = outputs.image_embeddings
        >>> text_embeddings = outputs.text_embeddings
        >>> multimodal_embeddings = outputs.multimodal_embeddings

        >>> outputs.image_embeddings.shape
        torch.Size([1, 197, 768])

        >>> text_embeddings.shape
        torch.Size([1, 7, 768])

        >>> multimodal_embeddings.shape
        torch.Size([1, 205, 768])
        ```
        NzRFLAVA model requires hidden states to work. Please set `output_hidden_states=True`)r“   r”   rî   rï   r6  r7  r   r|   r{   )rÓ   rî   rÁ   rÄ   rï   r6  r7  r   ©rÕ   r�   )rî   r7  )r   r    r!   r"   r#   r$   )ra   r7  r¶   r–  r›  r•  rœ  rƒ   r—  rX  r9   ÚonesrÕ   r‹   r   )r-   rÓ   r“   rî   rÄ   r”   rÁ   r£  r¤  rï   r6  r7  rw  r   Úimage_statesÚimage_mm_projectionr    r!   Útext_statesÚtext_mm_projectionr"   r#   r$   r™   rš   r›   Úattention_mask_imageÚattention_multimodalÚmultimodal_inputs                                r.   rŸ   zFlavaModel.forward:  s?  € ðL &1Ð%<�k�kÀ$Ä+ÔBYˆØ#ð 	sÝÐqÑrÔrÐrØÐØˆØ"ÐØˆØÐ#Ø×+Ò+Ø)Ø /Ø3Ø"3Ø%9Ø'ð ,ñ ô ˆLð .:¸!¬_¸lÈ1¼o˜lÐà"&×"=Ò"=¸lÈ2Ô>NÑ"OÔ"OÐàˆØˆØ!ÐØˆØÐ ØŸ/š/Ø#Ø-Ø)Ø-Ø"3Ø%9Ø'ð *ñ ô ˆKð ,7°q¬>¸;Àq¼>˜[ˆOà!%×!;Ò!;¸KÈ¼OÑ!LÔ!LÐà $ÐØ ÐØÐ*Ð/AÐ/MÐVmÐ/MØÐ)Ø)<Ô)BÑ&�
˜G QØÔ(Ô6ð !Ø˜q‘L�GÝ',¤z°*¸gÐNaÔNhÐ'iÑ'iÔ'iÐ$Ý',¤yÐ2FÈÐ1WÐ]^Ð'_Ñ'_Ô'_Ð$Ð$à'+Ð$Ý$œyÐ*=Ð?QÐ)RÐXYÐZÑZÔZÐØ $× 5Ò 5Ø Ð1EÐS^ð !6ñ !ô !Ðð %6°aÔ$8Ð!àð 	à ØØØØ%Ø!ðð õ  Ø-Ø%Ø+Ø#Ø"7Ø/ð
ñ 
ô 
ð 	
r0   rÜ   )NNNNNNNNNTN)r5   r6   r7   r   r;   ri   r   r   r9   r¢   r   r   r2   r   rŸ  r¤   rI   r¢  Ú
LongTensorr:   r   rŸ   r¥   r¦   s   @r.   rY  rY  ª  sœ  ø€ € € € € € àÐÐÑð)˜{ð )ð )ð )ð )ð )ð )ðV Øð /3Ø.2Ø,0ð/ð /à”<ð/ð œ tÑ+ð/ð œ tÑ+ð	/ð
 ”l TÑ)ð/ð Ð+Ô,ð/ð 
Ð+Ñ	+ð/ð /ð /ñ „^ñ Ôð/ðb Øð 48Ø04Ø.2ð*ð *à”lð*ð Ô)¨DÑ0ð*ð #'¨¡+ð	*ð
 œ tÑ+ð*ð Ð+Ô,ð*ð 
Ð+Ñ	+ð*ð *ð *ñ „^ñ Ôð*ðX ð .2Ø15Ø.2Ø.2Ø/3Ø04Ø48Ø/3Ø)-Ø%)Ø#'ðN
ð N
àÔ# dÑ*ðN
ð Ô'¨$Ñ.ðN
ð œ tÑ+ð	N
ð
 œ tÑ+ðN
ð œ¨Ñ,ðN
ð Ô&¨Ñ-ðN
ð $œl¨TÑ1ðN
ð "&¨¡ðN
ð   $™;ðN
ð #ðN
ð ˜D‘[ðN
ð 
Ð!Ñ	!ðN
ð N
ð N
ñ „^ðN
ð N
ð N
ð N
ð N
r0   rY  c                   óL   ‡ — e Zd Zdedefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚFlavaImageCodebookResPathÚin_sizeÚout_sizec                 ó(  •— t          ¦   «                              ¦   «          |dz  }t          ¦   «         }t          j        ¦   «         |d<   t          j        ||dd¬¦  «        |d<   t          j        ¦   «         |d<   t          j        ||dd¬¦  «        |d<   t          j        ¦   «         |d	<   t          j        ||dd¬¦  «        |d
<   t          j        ¦   «         |d<   t          j        ||dd¬¦  «        |d<   t          j        |¦  «        | _        d S )Né   Úrelu_1r   r   ©r¬   ÚpaddingÚconv_1Úrelu_2Úconv_2Úrelu_3Úconv_3Úrelu_4r   Úconv_4)rh   ri   r   r   ÚReLUr²   Ú
SequentialÚpath)r-   r²  r³  rw  Úhid_sizerÂ  rv   s         €r.   ri   z"FlavaImageCodebookResPath.__init__Í  sì   ø€ Ý‰Œ×ÒÑÔÐØ˜q‘=ˆå‰}Œ}ˆÝœ™œˆˆX‰Ýœ 7¨HÀ!ÈQÐOÑOÔOˆˆX‰Ýœ™œˆˆX‰Ýœ 8¨XÀ1ÈaÐPÑPÔPˆˆX‰Ýœ™œˆˆX‰Ýœ 8¨XÀ1ÈaÐPÑPÔPˆˆX‰Ýœ™œˆˆX‰Ýœ 8¨XÀ1ÈaÐPÑPÔPˆˆX‰å”M $Ñ'Ô'ˆŒ	ˆ	ˆ	r0   r¹   r%   c                 ó,   — |                       |¦  «        S r]   )rÂ  ©r-   r¹   s     r.   rŸ   z!FlavaImageCodebookResPath.forwardÝ  s   € Ø�yŠy˜‰|Œ|Ðr0   ©	r5   r6   r7   r£   ri   r9   r¢   rŸ   r¥   r¦   s   @r.   r±  r±  Ì  sq   ø€ € € € € ð( ð (¨sð (ð (ð (ð (ð (ð (ð ˜œð ¨%¬,ð ð ð ð ð ð ð ð r0   r±  c                   óP   ‡ — e Zd Zdededefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚFlavaImageCodebookBlockr²  r³  Ú
num_layersc                 ó  •— t          ¦   «                              ¦   «          d|dz  z  | _        ||k    rt          j        ||dd¬¦  «        | _        nt          j        ¦   «         | _        t          ||¦  «        | _        d S )Nr   r|   r   r·  )	rh   ri   Ú	post_gainr   r²   Úid_pathÚIdentityr±  Úres_path)r-   r²  r³  rÉ  rw  rv   s        €r.   ri   z FlavaImageCodebookBlock.__init__â  sr   ø€ Ý‰Œ×ÒÑÔÐà˜j¨!™mÑ,ˆŒà�hÒÐÝœ9 W¨hÀAÈqÐQÑQÔQˆDŒLˆLåœ;™=œ=ˆDŒLå1°'¸8ÑDÔDˆŒˆˆr0   r¹   r%   c                 óh   — |                       |¦  «        | j        |                      |¦  «        z  z   S r]   )rÌ  rË  rÎ  rÅ  s     r.   rŸ   zFlavaImageCodebookBlock.forwardî  s*   € Ø�|Š|˜A‰Œ ¤°$·-²-ÀÑ2BÔ2BÑ!BÑBÐBr0   rÆ  r¦   s   @r.   rÈ  rÈ  á  s‹   ø€ € € € € ð
E ð 
E¨sð 
EÀð 
Eð 
Eð 
Eð 
Eð 
Eð 
EðC˜œð C¨%¬,ð Cð Cð Cð Cð Cð Cð Cð Cr0   rÈ  c                   óZ   ‡ — e Zd Zddededededef
ˆ fd„Zdej        d	ej        fd
„Zˆ xZ	S )ÚFlavaImageCodebookLayerGroupTÚ
num_blocksrÉ  r²  r³  Úuse_poolc                 ód  •— t          ¦   «                              ¦   «          t          ¦   «         }t          |¦  «        D ]=}|dk    rt	          |||¦  «        |d|dz   › �<   Œ#t	          |||¦  «        |d|dz   › �<   Œ>|rt          j        d¬¦  «        |d<   t          j        |¦  «        | _        d S )Nr   Úblock_r   r|   )r¬   Úpool)	rh   ri   r   r2  rÈ  r   Ú	MaxPool2drÁ  Úgroup)	r-   rÒ  rÉ  r²  r³  rÓ  Úblocksr?  rv   s	           €r.   ri   z%FlavaImageCodebookLayerGroup.__init__ó  sÅ   ø€ Ý‰Œ×ÒÑÔÐÝ‘”ˆÝ�zÑ"Ô"ð 	cð 	cˆAØ�AŠvˆvÝ+BÀ7ÈHÐV`Ñ+aÔ+a�Ð'  A¡Ð'Ð'Ñ(Ð(å+BÀ8ÈXÐWaÑ+bÔ+b�Ð'  A¡Ð'Ð'Ñ(Ð(àð 	9Ýœ\°aÐ8Ñ8Ô8ˆF�6‰Nå”] 6Ñ*Ô*ˆŒ
ˆ
ˆ
r0   r¹   r%   c                 ó,   — |                       |¦  «        S r]   )rØ  rÅ  s     r.   rŸ   z$FlavaImageCodebookLayerGroup.forward  s   € Ø�zŠz˜!‰}Œ}Ðr0   r{  )
r5   r6   r7   r£   rI   ri   r9   r¢   rŸ   r¥   r¦   s   @r.   rÑ  rÑ  ò  s�   ø€ € € € € ð+ð + 3ð +°Cð +À#ð +ÐQTð +Ð`dð +ð +ð +ð +ð +ð +ð˜œð ¨%¬,ð ð ð ð ð ð ð ð r0   rÑ  a"  
    The FLAVA's image codebook model inspired from DALL-E's original encoder. Outputs raw hidden states and can be used
    to generate image tokens for an image based on DALL-E's vocab. Used to generate labels for MIM. Use
    `get_codebook_indices` to get image tokens for an image.
    c                   ó°   ‡ — e Zd ZU dZeed<   dZdZdZdede	fˆ fd„Z
dej        dej        fd	„Zdej        dej        fd
„Zdej        dej        fd„Zˆ xZS )ÚFlavaImageCodebookÚmodelra   r“   rc  Frw  c                 ó&  •— t          ¦   «                              |¦  «         || _        |j        | _        |j        | _        |j        | _        |j        | _        |j        | _        | j        | j        z  }t          ¦   «         }t          j
        ¦   «         |d<   t          j        d| j        z  | j        dd¬¦  «        |d<   t          ¦   «         }t          j        | j        d| j        z  dd¬¦  «        |d	<   t          | j        |d| j        z  d| j        z  ¦  «        |d
<   t          | j        |d| j        z  d| j        z  ¦  «        |d<   t          | j        |d| j        z  d| j        z  ¦  «        |d<   t          | j        |d| j        z  d| j        z  d¬¦  «        |d<   t          j        |¦  «        |d<   t          j        |¦  «        | _        |                      ¦   «          | j        j        r|                      ¦   «         D ]}d|_        Œ
d S d S )NÚrelué   r   r   r·  Úconvé   r   ÚinputÚgroup_1r|   Úgroup_2rµ  Úgroup_3F)rÓ  Úgroup_4r  )rh   ri   ra   Ú
num_groupsÚinput_channelsÚnum_blocks_per_grouprm   rÇ   r   r   rÀ  r²   rÑ  rÁ  rÙ  rk  ÚfreezeÚ
parametersÚrequires_grad)r-   ra   rw  rÉ  Úoutput_blocksrÙ  Úparamrv   s          €r.   ri   zFlavaImageCodebook.__init__  s&  ø€ õ
 	‰Œ×Ò˜Ñ Ô Ð àˆŒØ Ô+ˆŒØ$Ô3ˆÔØ$*Ô$?ˆÔ!Ø!Ô-ˆÔØ Ô+ˆŒà”_ tÔ'@Ñ@ˆ
å#™œˆÝ "¤¡	¤	ˆ�fÑÝ "¤	¨!¨dÔ.>Ñ*>ÀÄÐ]^ÐhiÐ jÑ jÔ jˆ�fÑå‘”ˆÝœ) DÔ$7¸¸TÔ=MÑ9MÐ[\ÐfgÐhÑhÔhˆˆw‰Ý8ØÔ% z°1°tÔ7GÑ3GÈÈTÔM]ÑI]ñ
ô 
ˆˆyÑõ 9ØÔ% z°1°tÔ7GÑ3GÈÈTÔM]ÑI]ñ
ô 
ˆˆyÑõ 9ØÔ% z°1°tÔ7GÑ3GÈÈTÔM]ÑI]ñ
ô 
ˆˆyÑõ 9ØÔ% z°1°tÔ7GÑ3GÈÈTÔM]ÑI]Ðhmð
ñ 
ô 
ˆˆyÑõ œ=¨Ñ7Ô7ˆˆxÑå”m FÑ+Ô+ˆŒà�ŠÑÔÐàŒ;Ôð 	,ØŸšÑ*Ô*ð ,ð ,�Ø&+�Ô#Ð#ð	,ð 	,ð,ð ,r0   r%   c                 ó~   — dt           › dt           › d� |                      |¦  «        }t          j        |d¬¦  «        S )NaI  
        Args:
            pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
                Pixel values. Codebook pixel values can be obtained using [`AutoImageProcessor`] by passing
                `return_codebook_pixels=True`. See [`FlavaImageProcessor.__call__`] for details.

        Examples:
        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoImageProcessor, FlavaImageCodebook

        >>> model = FlavaImageCodebook.from_pretrained("úE")
        >>> image_processor = AutoImageProcessor.from_pretrained("a¹  ")

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))

        >>> inputs = image_processor([image], return_codebook_pixels=True, return_tensors="pt")
        >>> inputs = dict(pixel_values=inputs.codebook_pixel_values)

        >>> outputs = model.get_codebook_indices(**inputs)
        ```
        r   )Úaxis)Ú_CHECKPOINT_FOR_CODEBOOK_DOCrÙ  r9   Úargmax©r-   r“   Úz_logitss      r.   Úget_codebook_indicesz'FlavaImageCodebook.get_codebook_indices@  sZ   € ð	õ :Vð	ð 	õ D`ð	ð 	ð 	ð 	ð4 —;’;˜|Ñ,Ô,ˆÝŒ|˜H¨1Ð-Ñ-Ô-Ð-r0   c                 óh   — |                       |¦  «        } t          j        d¬¦  «        |¦  «        S )Nr   r�   )rÙ  r   ÚSoftmaxrõ  s      r.   Úget_codebook_probsz%FlavaImageCodebook.get_codebook_probs^  s0   € Ø—;’;˜|Ñ,Ô,ˆØ �rŒz˜aÐ Ñ Ô  Ñ*Ô*Ð*r0   c                 ó(  — dt           › dt           › d� t          |j        ¦  «        dk    rt          d|j        › d�¦  «        ‚|j        d         | j        k    r%t          d|j        d         › d	| j        › �¦  «        ‚|                      |¦  «        S )
NaJ  
        Args:
            pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
                Pixel values. Codebook pixel values can be obtained using [`AutoImageProcessor`] by passing
                `return_codebook_pixels=True`. See [`FlavaImageProcessor.__call__`] for details.

        Examples:

        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoImageProcessor, FlavaImageCodebook

        >>> model = FlavaImageCodebook.from_pretrained("rñ  aÖ  ")

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))

        >>> inputs = image_processor([image], return_codebook_pixels=True, return_tensors="pt")
        >>> inputs = dict(pixel_values=inputs.codebook_pixel_values)

        >>> outputs = model(**inputs)
        >>> print(outputs.shape)
        (1, 196)
        ```
        rµ  zinput shape z
 is not 4dr   z
input has z channels but model built for )ró  Úlenrƒ   r¶   ré  rÙ  )r-   r“   rw  s      r.   rŸ   zFlavaImageCodebook.forwardb  sº   € ð	õ :Vð	ð 	õ D`ð	ð 	ð 	ð 	õ: ˆ|Ô!Ñ"Ô" aÒ'Ð'ÝÐJ¨LÔ,>ÐJÐJÐJÑKÔKÐKØÔ˜aÔ  DÔ$7Ò7Ð7ÝÐt¨,Ô*<¸QÔ*?ÐtÐtÐ_cÔ_rÐtÐtÑuÔuÐuØ�{Š{˜<Ñ(Ô(Ð(r0   )r5   r6   r7   r]  r   r;   r}  r^  r_  r   ri   r9   r¢   r÷  rú  r:   rŸ   r¥   r¦   s   @r.   rÜ  rÜ    sê   ø€ € € € € € ð  ÐØ$Ð$Ð$Ñ$Ø$€OØ!ÐØ&+Ð#ð*,à(ð*,ð ð*,ð *,ð *,ð *,ð *,ð *,ðX.°´ð .À%Ä,ð .ð .ð .ð .ð<+¨u¬|ð +ÀÄð +ð +ð +ð +ð") EÔ$5ð ")ÀEÄLð ")ð ")ð ")ð ")ð ")ð ")ð ")ð ")r0   rÜ  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFlavaPredictionHeadTransformc                 óV  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          |j        t          ¦  «        rt          |j                 | _
        n|j        | _
        t          j        |j        |j        ¬¦  «        | _        d S )Nr¿   )rh   ri   r   rç   rm   r  r®   r  r  r	   Útransform_act_fnrÍ   rÎ   rÒ   s     €r.   ri   z%FlavaPredictionHeadTransform.__init__ˆ  s…   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3EÑFÔFˆŒ
Ý�fÔ'­Ñ-Ô-ð 	6Ý$*¨6Ô+<Ô$=ˆDÔ!Ð!à$*Ô$5ˆDÔ!Ýœ fÔ&8¸fÔ>SÐTÑTÔTˆŒˆˆr0   c                 ó„   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|S r]   )r  r   rÍ   r  s     r.   rŸ   z$FlavaPredictionHeadTransform.forward‘  s=   € ØŸ
š
 =Ñ1Ô1ˆØ×-Ò-¨mÑ<Ô<ˆØŸš }Ñ5Ô5ˆØÐr0   ©r5   r6   r7   ri   rŸ   r¥   r¦   s   @r.   rþ  rþ  ‡  sL   ø€ € € € € ðUð Uð Uð Uð Uðð ð ð ð ð ð r0   rþ  c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )rS  Nc                 óJ  •— t          ¦   «                              ¦   «          || _        t          |¦  «        | _        t          j        |j        |j        d¬¦  «        | _	        t          j
        t          j        |j        ¦  «        ¦  «        | _        |�|| j	        _        d S d S )NTrâ   )rh   ri   ra   rþ  Ú	transformr   rç   rm   rÇ   Údecoderrk   r9   rl   rã   Úweight)r-   ra   r  rv   s      €r.   ri   z"FlavaMaskedPredictionHead.__init__™  s‰   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝ5°fÑ=Ô=ˆŒÝ”y Ô!3°VÔ5FÈTÐRÑRÔRˆŒÝ”L¥¤¨VÔ->Ñ!?Ô!?Ñ@Ô@ˆŒ	ØÐØ"(ˆDŒLÔÐÐð Ðr0   c                 óZ   — |                       |¦  «        }|                      |¦  «        }|S r]   )r  r  rÅ  s     r.   rŸ   z!FlavaMaskedPredictionHead.forward¢  s'   € Ø�NŠN˜1ÑÔˆØ�LŠL˜‰OŒOˆØˆr0   r]   r  r¦   s   @r.   rS  rS  ˜  sL   ø€ € € € € ð)ð )ð )ð )ð )ð )ðð ð ð ð ð ð r0   rS  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFlavaITMHeadc                 ó¼   •— t          ¦   «                              ¦   «          || _        t          |¦  «        | _        t          j        |j        d¦  «        | _        d S )Nr|   )	rh   ri   ra   rC  rj  r   rç   rm   Úseq_relationshiprÒ   s     €r.   ri   zFlavaITMHead.__init__©  sL   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝ! &Ñ)Ô)ˆŒÝ "¤	¨&Ô*<¸aÑ @Ô @ˆÔÐÐr0   c                 óZ   — |                       |¦  «        }|                      |¦  «        }|S r]   )rj  r  rÅ  s     r.   rŸ   zFlavaITMHead.forward¯  s)   € Ø�KŠK˜‰NŒNˆØ×!Ò! !Ñ$Ô$ˆØˆr0   r  r¦   s   @r.   r
  r
  ¨  sL   ø€ € € € € ðAð Að Að Að Aðð ð ð ð ð ð r0   r
  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFlavaGlobalContrastiveHeadc                 ón   •— t          ¦   «                              ¦   «          || _        |j        | _        d S r]   )rh   ri   ra   Úglobal_backprop_contrastiverÒ   s     €r.   ri   z#FlavaGlobalContrastiveHead.__init__¶  s1   ø€ Ý‰Œ×ÒÑÔÐØˆŒØ+1Ô+MˆÔ(Ð(Ð(r0   c                 óœ  ‡‡— t          j        |¦  «        }t           j                             ¦   «         rt           j                             ¦   «         s6t          j        ‰                     d¦  «        ‰j        ¬¦  «        }‰g}‰g}�n@‰                     d¦  «        }t           j                             ¦   «         }	| j	        rSt           j        j
        j                             ‰¦  «        }t           j        j
        j                             ‰¦  «        }nvˆfd„t          |	¦  «        D ¦   «         }ˆfd„t          |	¦  «        D ¦   «         }t           j                             |‰¦  «         t           j                             |‰¦  «         |t           j                             ¦   «         z  t          j        |‰j        ¬¦  «        z   }t          j        |¦  «        }t          j        |¦  «        }t          j        ‰|                     dd¦  «        ¦  «        |z  }
t          j        ‰|                     dd¦  «        ¦  «        |z  }|
||fS )Nr   r¦  c                 ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS r<   ©r9   Ú
zeros_like)r+   r›   r!   s     €r.   r0  z6FlavaGlobalContrastiveHead.forward.<locals>.<listcomp>Ë  s$   ø€ Ð'eÐ'eÐ'eÈa­Ô(8¸Ñ(IÔ(IÐ'eÐ'eÐ'er0   c                 ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS r<   r  )r+   r›   r   s     €r.   r0  z6FlavaGlobalContrastiveHead.forward.<locals>.<listcomp>Ì  s%   ø€ Ð&eÐ&eÐ&eÈa¥uÔ'7Ð8HÑ'IÔ'IÐ&eÐ&eÐ&er0   r   )r9   ÚexpÚdistributedÚis_availableÚis_initializedrÐ   r~   rÕ   Úget_world_sizer  r   rˆ   Ú
all_gatherr2  Úget_rankr‹   rò   r¸   )r-   r   r!   r[  ÚtemperatureÚlabelsÚimage_embeddings_allÚtext_embeddings_allÚlocal_batch_sizeÚ
world_sizeÚlogits_per_imageÚlogits_per_texts    ``         r.   rŸ   z"FlavaGlobalContrastiveHead.forward»  s2  øø€ Ý”i Ñ,Ô,ˆÝÔ ×-Ò-Ñ/Ô/ð 	µuÔ7H×7WÒ7WÑ7YÔ7Yð 	Ý”\Ð"2×"7Ò"7¸Ñ":Ô":ÐCSÔCZÐ[Ñ[Ô[ˆFØ$4Ð#5Ð Ø#2Ð"3ÐÑà/×4Ò4°QÑ7Ô7ÐÝÔ*×9Ò9Ñ;Ô;ˆJàÔ/ð 	Sõ (-Ô'8Ô';Ô'F×'QÒ'QÐRbÑ'cÔ'cÐ$Ý&+Ô&7Ô&:Ô&E×&PÒ&PÐQ`Ñ&aÔ&aÐ#Ð#à'eÐ'eÐ'eÐ'eÕSXÐYcÑSdÔSdÐ'eÑ'eÔ'eÐ$Ø&eÐ&eÐ&eÐ&eÕSXÐYcÑSdÔSdÐ&eÑ&eÔ&eÐ#ÝÔ!×,Ò,Ð-AÐCSÑTÔTÐTÝÔ!×,Ò,Ð-@À/ÑRÔRÐRà%­Ô(9×(BÒ(BÑ(DÔ(DÑDÅuÄ|Ø Ð)9Ô)@ðHñ Hô Hñ ˆFõ  %œyÐ)=Ñ>Ô>ÐÝ#œiÐ(;Ñ<Ô<Ðå œ<Ð(8Ð:M×:WÒ:WÐXYÐ[\Ñ:]Ô:]Ñ^Ô^ÐalÑlÐÝœ, Ð8L×8VÒ8VÐWXÐZ[Ñ8\Ô8\Ñ]Ô]Ð`kÑkˆà °&Ð8Ð8r0   r  r¦   s   @r.   r  r  µ  sL   ø€ € € € € ðNð Nð Nð Nð Nð
9ð 9ð 9ð 9ð 9ð 9ð 9r0   r  zk
    The FLAVA model for pretraining which outputs losses, embeddings, logits and transformer outputs.
    c            '       óÖ  ‡ — e Zd ZdddddœZd dedej        dz  fˆ fd	„Zd
ej	        fd„Z
e	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d!dej        dz  dej        dz  dej        dz  dej        dz  dej	        dz  dej	        dz  dej	        dz  dej        dz  dej	        dz  dedz  dej	        dz  dej	        dz  dej	        dz  dedz  dededz  dedz  deej	                 ez  f$d„¦   «         Zˆ xZS )"ÚFlavaForPreTrainingzmmm_text_head.decoder.biaszmim_head.decoder.biaszmlm_head.decoder.biaszmmm_image_head.decoder.bias)zmmm_text_head.biaszmim_head.biaszmlm_head.biaszmmm_image_head.biasNra   Úimage_codebookc                 ó  •— t          ¦   «                              |¦  «         t          |¦  «        | _        || _        | j        € |j        rt          |j        ¦  «        | _        t          |j	        ¦  «        | _
        t          |j        ¦  «        | _        t          |¦  «        | _        t          |j	        ¦  «        | _        t          |j        ¦  «        | _        t#          |¦  «        | _        |j	        j        | _        |j        j        | _        |j        | _        |j        | _        |j        | _        |j        | _        |j        | _        |j        | _        |j        | _        |j        | _        |                      ¦   «          dS )zò
        image_codebook ([`nn.Module`]):
            If passed, the image codebook will be set to this. Otherwise, it will be initialized using the
            image_codebook_config defined in the config first as the first parameter.
        N)rh   ri   rY  rM  r(  Úinit_codebookrÜ  Úimage_codebook_configrS  r�  Úmim_headrŒ  Úmlm_headr
  Úitm_headÚmmm_image_headÚmmm_text_headr  Úglobal_contrastive_headrÇ   Úimage_vocab_sizeÚtext_vocab_sizeÚ
mlm_weightÚ
mim_weightÚglobal_contrastive_weightÚce_ignore_indexÚ
itm_weightÚmmm_image_weightÚmmm_text_weightÚ skip_unmasked_multimodal_encoderrk  )r-   ra   r(  rv   s      €r.   ri   zFlavaForPreTraining.__init__ë  sJ  ø€ õ 	‰Œ×Ò˜Ñ Ô Ð Ý Ñ'Ô'ˆŒ
à,ˆÔØÔÐ&¨6Ô+?Ð&Ý"4°VÔ5QÑ"RÔ"RˆDÔõ 2°&Ô2EÑFÔFˆŒÝ1°&Ô2DÑEÔEˆŒÝ$ VÑ,Ô,ˆŒÝ7¸Ô8KÑLÔLˆÔÝ6°vÔ7IÑJÔJˆÔÝ'AÀ&Ñ'IÔ'IˆÔ$à &Ô 3Ô >ˆÔØ%Ô1Ô<ˆÔØ Ô+ˆŒØ Ô+ˆŒØ)/Ô)IˆÔ&Ø%Ô5ˆÔØ Ô+ˆŒØ &Ô 7ˆÔØ%Ô5ˆÔØ06Ô0WˆÔ-à�ŠÑÔÐÐÐr0   r¹   c                 óˆ   — |                      ¦   «         dk    r)|                     |                     d¦  «        d¦  «        }|S )Nr|   r   r{   )r‚   rŠ   r~   rÅ  s     r.   Ú_resize_to_2dz!FlavaForPreTraining._resize_to_2d  s5   € Ø�5Š5‰7Œ7�QŠ;ˆ;Ø—’�q—v’v˜a‘y”y "Ñ%Ô%ˆAØˆr0   TrÓ   Úinput_ids_maskedr“   Úcodebook_pixel_valuesrî   rÄ   r”   rÁ   r£  r;  Ú
mlm_labelsÚ
mim_labelsÚ
itm_labelsrï   r6  r7  Úreturn_lossr%   c                 ó¶  — |�|n| j         j        }|�|n| j         j        }|
�|
n| j        }
|€|�t                               d¦  «         |}|                      ||||||	|
||d¬¦
  «
        }|                      |||||	|||d¬¦	  «	        }d}|j        }|j        }|j        }|j        }|j	        }dx}x}x}x}x}x} }!dx}"x}#x}$}%dx}&x}'}(|€|�E|€C|rA| j
        €t          d¦  «        ‚|€t          d¦  «        ‚| j
                             |¦  «        }| j        dk    �r(|��%|�€"|})|��|                      |¦  «        }|                      |¦  «        }| j        ||                     d¦  «        <   |)dd…|                     d	¦  «         d…dd…f         })|                     | j        ¦  «        }*||*         }+|)|*dd…f         })|                      |)¦  «        }"|rVt(          j                             |"                     d
| j        ¦  «        |+                     d
¦  «        ¦  «        }|| j        z  }n|                      |)¦  «        }"| j        dk    ró|�ñ|€ï|},|�Ö|                      |¦  «        }|,dd…|                     d	¦  «         d…dd…f         },|                     | j        ¦  «        }*||*         }-|,|*dd…f         },|                      |,¦  «        }#|rVt(          j                             |#                     d
| j        ¦  «        |-                     d
¦  «        ¦  «        }|| j        z  }n|                      |,¦  «        }#| j        dk    r˜|�–|                      |¦  «        }&|�|                     d¦  «        }.|.|.                     ¦   «          z  }|r*t(          j                             |&|¦  «        }!|!| j        z  }!|�||         }|�||         }|�||         }||         }|��4| j        dk    �r(|})|                     d	¦  «        d	z
  }/|)dd…dd|/z   …dd…f         })|�ã|                      |¦  «        }|                      |¦  «        }| j        ||                     d¦  «        <   |                     | j        ¦  «        }*||*         }+|)|*dd…f         })|                       |)¦  «        }%|rVt(          j                             |%                     d
| j        ¦  «        |+                     d
¦  «        ¦  «        }|| j        z  }n|                       |)¦  «        }%|�ú| j!        dk    rï|},|,dd…|                     d	¦  «         d…dd…f         },|�±|                      |¦  «        }|                     | j        ¦  «        }*||*         }-|,|*dd…f         },|  "                    |,¦  «        }$|rVt(          j                             |$                     d
| j        ¦  «        |-                     d
¦  «        ¦  «        }|| j!        z  }n|  "                    |,¦  «        }$|��h|��e| j#        dk    �rY| j         $                    |dd…ddd…f         ¦  «        }0t(          j         %                    |0d
¬¦  «        }0| j         &                    |dd…ddd…f         ¦  «        }1t(          j         %                    |1d
¬¦  «        }1| j'        r/| j        j(        j)         *                    tV          tX          ¦  «         |  -                    |1|0| j        j(        ¦  «        \  }'}(}2|�|'|         }'|(|         }(|2|         }2|rRt(          j                             |'|2¦  «        }3t(          j                             |(|2¦  «        }4|3|4z   dz  } | | j#        z  } t]          |||!| ||¬¦  «        }5|r?|5 /                    ¦   «         s+ta          d„ |5 1                    ¦   «         D ¦   «         ¦  «        }|�s||j2        �|j2         3                    ¦   «         nd||j4        �|j4         3                    ¦   «         nd|j	        |j5        �|j5         3                    ¦   «         nd||j2        �|j2         3                    ¦   «         nd||j4        �|j4         3                    ¦   «         nd||j5        �|j5         3                    ¦   «         nd|"|#|&|'|(|%|$f}6|r|5 /                    ¦   «         s||5f|6z   }6tm          d„ |6D ¦   «         ¦  «        S to          d%i d|“d|5“d|“d|j2        “d|“d|j4        “d|j	        “d|j5        “d|“d|j2        “d|“d|j4        “d|“d|j5        “d|"“d|#“d |&“d!|'“d"|(“d#|%“d$|$“ŽS )&aè  
        input_ids (`torch.LongTensor` of shape `(batch_size, text_seq_len)`):
            Indices of input sequence tokens in the vocabulary. Indices can be obtained using [`AutoTokenizer`]. See
            [`PreTrainedTokenizer.encode`] and [`PreTrainedTokenizer.__call__`] for details. [What are input
            IDs?](../glossary#input-ids)
        input_ids_masked (`torch.LongTensor` of shape `(batch_size, text_seq_len)`):
            Indices of input sequence tokens in the vocabulary. These ones are the masked version of the original task
            to be used with MLM. Indices can be obtained using [`AutoTokenizer`] along with
            [`DataCollatorForMaskedLanguageModeling`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details. [What are input IDs?](../glossary#input-ids)
        codebook_pixel_values (`torch.FloatTensor` of shape `(batch_size, num_image_patches, patch_size, patch_size, 3)`, *optional*):
            Pixel values for image patches that are used to compute the image codebook labels for masked image modeling.
        token_type_ids (`torch.LongTensor` of shape `(batch_size, text_seq_len)`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:
            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.
            [What are token type IDs?](../glossary#token-type-ids)
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, image_num_patches)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).
        image_attention_mask (`torch.FloatTensor` of shape `(batch_size, image_num_patches)`, *optional*):
            Mask to avoid performing attention on padding token indices specifically for images. 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)
        skip_unmasked_multimodal_encoder (*bool*, *optional*):
            Skip any calculations for multimodal encoder for unmasked inputs. FLAVA pretraining doesn't need unmasked
            multimodal embeddings or outputs as of now.
        mlm_labels (`torch.LongTensor` of shape `(batch_size, text_seq_len)`, *optional*):
            Labels for computing the left-to-right language and multimodal masked modeling loss (next word prediction).
            Indices should be in `[-100, 0, ..., text_config.vocab_size - 1]` (see `input_ids` docstring). Tokens with
            indices set to `-100` are ignored (masked), the loss is only computed for the tokens with labels in `[0,
            ..., text_config.vocab_size - 1]`.
        mim_labels (`torch.LongTensor` of shape `(batch_size, image_num_patches)`, *optional*):
            Labels for computing the image and multimodal masked modeling loss. Indices should be in `[-100, 0, ...,
            image_config.vocab_size - 1]`. Tokens with indices set to `-100` are ignored (masked), the loss is only
            computed for the tokens with labels in `[0, ..., image_config.vocab_size - 1]`. If not passed, they are
            generated automatically using the image codebook assigned to the model. By default, it uses
            [`FlavaImageCodebook`]. See [`FlavaImageCodebook`] to understand how to generate mim_labels.
        itm_labels (`torch.LongTensor` of shape `(batch_size, 1)`, *optional*):
            Labels for computing the image-text matching loss. 0 means the pairs don't match and 1 means they match.
            The pairs with 0 will be skipped for calculation of MMM and global contrastive losses as well.
        return_loss (`bool`, *optional*, default to None):
            Whether to return calculated loss or not.

        Examples:
        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import FlavaForPreTraining, AutoProcessor

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))

        >>> model = FlavaForPreTraining.from_pretrained("facebook/flava-full")
        >>> processor = AutoProcessor.from_pretrained("facebook/flava-full")

        >>> text = ["a photo of a cat"]

        >>> inputs = processor(
        ...     images=[image],
        ...     text=text,
        ...     return_masks=True,
        ...     return_codebook_pixels=True,
        ...     padding=True,
        ...     max_length=77,
        ...     return_tensors="pt",
        ... )


        >>> output = model(**inputs)
        ```
        Nzð`input_ids_masked` isn't passed which means MLM loss won't be calculated correctlySetting it to `input_ids` so that model can work. Please pass it if this is unintentional. This is usually OKAY if you are doing inference on unmasked text...T)
rÓ   r“   rî   rÄ   rÁ   r£  r¤  rï   r6  r7  )	rÓ   r“   rî   rÄ   r£  r”   rï   r6  r7  zÊ`return_loss` is set to True but the image codebook is not initialized and no `mim_labels`  have been passed. Reinstantiate the model with `init_codebook` set to True or pass in your custom `mim_labels`z�`codebook_pixel_value` are required to generate `mim_labels` if loss is expected. Call `AutoProcessor` with `return_codebook_pixels` set to Truer   r   r{   r|   r�   )r?   r@   rA   rB   rC   rD   c              3   ó"   K  — | ]
}|�|ndV — Œd S rH  r<   )r+   rL   s     r.   r/   z.FlavaForPreTraining.forward.<locals>.<genexpr>K  s+   è è € Ð_Ð_À TÐ%5˜T˜T¸1Ð_Ð_Ð_Ð_Ð_Ð_r0   c              3   ó   K  — | ]}|®|V — Œ	d S r]   r<   )r+   r¹   s     r.   r/   z.FlavaForPreTraining.forward.<locals>.<genexpr>l  s"   è è € Ð<Ð<˜q¨a¨m˜¨m¨m¨m¨mÐ<Ð<r0   rL   rM   r   r    r!   r"   r#   r$   rN   rO   rP   rQ   rR   rS   rT   rU   rV   rW   rX   rY   rZ   r<   )8ra   r7  rC  r;  ÚloggerÚwarningrM  r   r!   r#   r(  ÚRuntimeErrorr¶   r÷  r5  r=  r7  Úner~   r,  r   rˆ   Úcross_entropyrŠ   r2  r4  r-  r3  r8  r.  Úanyr9  r/  r:  r0  r6  r™  Ú	normalizer˜  Útrainingr[  ÚdataÚclamp_ÚLOGIT_SCALE_CLAMP_MINÚLOGIT_SCALE_CLAMP_MAXr1  r>   rG   ÚsumrF   r    r*   r"   r$   r2   rK   )7r-   rÓ   r>  r“   r?  rî   rÄ   r”   rÁ   r£  r;  r@  rA  rB  rï   r6  r7  rC  rw  Úflava_outputÚflava_masked_outputÚpos_maskr   r!   rN   rP   rR   Ú
total_lossÚmim_lossÚmlm_lossÚmmm_text_lossÚmmm_image_lossÚgc_lossÚitm_lossrT   rU   rZ   rY   rV   r$  r%  Úsequence_for_imageÚmasked_tokensÚmim_labels_filteredÚsequence_for_textÚmlm_labels_filteredÚ	pos_pairsÚ	end_indexÚtext_embeddingÚimage_embeddingÚ	gc_labelsÚgc_loss_imageÚgc_loss_textÚflava_lossesr  s7                                                          r.   rŸ   zFlavaForPreTraining.forward  s¸  € ðD &1Ð%<�k�kÀ$Ä+ÔBYˆØ%0Ð%<�k�kÀ$Ä+ÔBYˆð 0Ð;ð -Ð,àÔ6ð 	)ð Ð#¨	Ð(=Ý�NŠNð?ñô ð ð
  )Ðà—z’zØØ%Ø)Ø)Ø%Ø!5ð %EØ/Ø!5àð "ñ 
ô 
ˆð  #ŸjšjØ&Ø%Ø)Ø)Ø!5Ø+Ø/Ø!5Øð )ñ 

ô 

Ðð ˆà'Ô8ÐØ&Ô6ˆØ"5Ô"FÐØ!4Ô!DÐØ':Ô'PÐ$àaeÐeˆ
Ðe�XÐe Ðe¨=Ðe¸>ÐeÈGÐV^ØGKÐKˆ
ÐK�ZÐK /Ð4DØ:>Ð>ˆ
Ð>Ð%¨ð #Ð.Ð2NÐ2ZØÐ! kÐ!ØÔ&Ð.Ý&ð;ñô ð ð
 )Ð0Ý$ðYñô ð ð "Ô0×EÒEÐF[Ñ\Ô\�
ð Œ?˜QÒÑÐ#:Ñ#FÐKgÑKoØ!8ÐàÑ%Ø!×/Ò/°
Ñ;Ô;�
Ø"&×"4Ò"4°_Ñ"EÔ"E�Ø7;Ô7K�
˜?×-Ò-¨dÑ3Ô3Ñ4à%7¸¸¸¸J¿OºOÈAÑ<NÔ<NÐ;NÐ;PÐ;PÐRSÐRSÐRSÐ8SÔ%TÐ"Ø *§¢¨dÔ.BÑ CÔ C�Ø&0°Ô&?Ð#Ø%7¸ÀqÀqÀqÐ8HÔ%IÐ"Ø!Ÿ]š]Ð+=Ñ>Ô>�
Øð 0Ý!œ}×:Ò:Ø"Ÿš¨¨DÔ,AÑBÔBÐDW×D\ÒD\Ð]_ÑD`ÔD`ñ ô  �Hð  ¤Ñ/�Høà!Ÿ]š]Ð+=Ñ>Ô>�
ð Œ?˜QÒÐÐ#9Ð#EÐJfÐJnØ 6ÐØÐ%Ø!×/Ò/°
Ñ;Ô;�
Ø$5°a°a°a¸*¿/º/È!Ñ:LÔ:LÐ9LÐ9NÐ9NÐPQÐPQÐPQÐ6QÔ$RÐ!Ø *§¢¨dÔ.BÑ CÔ C�Ø&0°Ô&?Ð#Ø$5°mÀQÀQÀQÐ6FÔ$GÐ!Ø!Ÿ]š]Ð+<Ñ=Ô=�
Øð 0Ý!œ}×:Ò:Ø"Ÿš¨¨DÔ,@ÑAÔAÐCV×C[ÒC[Ð\^ÑC_ÔC_ñ ô  �Hð  ¤Ñ/�Høà!Ÿ]š]Ð+<Ñ=Ô=�
ð Œ?˜QÒÐÐ#?Ð#KØŸšÐ'CÑDÔDˆJàÐ%Ø&ŸMšM¨!Ñ,Ô,�	à$¨	¯ª©¬Ð'7Ñ7�Øð 0Ý!œ}×:Ò:¸:ÀzÑRÔR�HØ ¤Ñ/�Hà/Ð;Ø3OÐPXÔ3YÐ0àÐ)Ø!+¨HÔ!5�JàÐ)Ø!+¨HÔ!5�JØ&5°hÔ&?�Oð (Ñ3¸Ô8MÐPQÒ8QÑ8QØ!=ÐØ/×4Ò4°QÑ7Ô7¸!Ñ;ˆIØ!3°A°A°A°q¸1¸y¹=Ð7HÈ!È!È!Ð4KÔ!LÐàÐ%Ø!×/Ò/°
Ñ;Ô;�
Ø"&×"4Ò"4°_Ñ"EÔ"E�Ø7;Ô7K�
˜?×-Ò-¨dÑ3Ô3Ñ4à *§¢¨dÔ.BÑ CÔ C�Ø&0°Ô&?Ð#Ø%7¸ÀqÀqÀqÐ8HÔ%IÐ"Ø#'×#6Ò#6Ð7IÑ#JÔ#JÐ Øð <Ý%'¤]×%@Ò%@Ø(×-Ò-¨b°$Ô2GÑHÔHÐJ]×JbÒJbÐceÑJfÔJfñ&ô &�Nð # dÔ&;Ñ;�Nøà#'×#6Ò#6Ð7IÑ#JÔ#JÐ ð (Ð3¸Ô8LÈqÒ8PÐ8PØ <ÐØ 1°!°!°!Ð6L×6QÒ6QÐRSÑ6TÔ6TÐ5TÐ5VÐ5VÐXYÐXYÐXYÐ2YÔ ZÐàÐ%Ø!×/Ò/°
Ñ;Ô;�
Ø *§¢¨dÔ.BÑ CÔ C�Ø&0°Ô&?Ð#Ø$5°mÀQÀQÀQÐ6FÔ$GÐ!Ø"&×"4Ò"4Ð5FÑ"GÔ"G�Øð :Ý$&¤M×$?Ò$?Ø'×,Ò,¨R°Ô1EÑFÔFÐH[×H`ÒH`ÐacÑHdÔHdñ%ô %�Mð " TÔ%9Ñ9�Møà"&×"4Ò"4Ð5FÑ"GÔ"G�ð Ñ'¨OÑ,GÈDÔLjÐmnÒLnÑLnØ!œZ×7Ò7¸ÈÈÈÈ1ÈaÈaÈaÈÔ8PÑQÔQˆNÝœ]×4Ò4°^ÈÐ4ÑLÔLˆNà"œj×9Ò9Ð:JÈ1È1È1ÈaÐQRÐQRÐQRÈ7Ô:SÑTÔTˆOÝ œm×5Ò5°oÈ2Ð5ÑNÔNˆOàŒ}ð aØ”
Ô&Ô+×2Ò2Õ3HÕJ_Ñ`Ô`Ð`à;?×;WÒ;WØ °´Ô1Gñ<ô <Ñ8Ð˜o¨yð
 Ð#Ø#3°HÔ#=Ð Ø"1°(Ô";�Ø% hÔ/�	àð :Ý "¤× ;Ò ;Ð<LÈiÑ XÔ X�Ý!œ}×:Ò:¸?ÈIÑVÔV�Ø(¨<Ñ7¸1Ñ<�Ø˜4Ô9Ñ9�å"ØØØØ&Ø$Ø"ð
ñ 
ô 
ˆð ð 	`˜|×4Ò4Ñ6Ô6ð 	`ÝÐ_Ð_È×I\ÒI\ÑI^ÔI^Ð_Ñ_Ô_Ñ_Ô_ˆJàñ 	=à Ø8DÔ8QÐ8]�Ô)×2Ò2Ñ4Ô4Ð4ÐcgØØ7CÔ7OÐ7[�Ô(×1Ò1Ñ3Ô3Ð3ÐaeØÔ2Ø=IÔ=[Ð=g�Ô.×7Ò7Ñ9Ô9Ð9ÐmqØ'Ø?RÔ?_Ð?kÐ#Ô0×9Ò9Ñ;Ô;Ð;ÐquØ&Ø>QÔ>]Ð>iÐ#Ô/×8Ò8Ñ:Ô:Ð:ÐosØ,à&Ô8ÐDð $Ô5×>Ò>Ñ@Ô@Ð@àØØØØ ØØ Øð+ˆFð. ð  <×#8Ò#8Ñ#:Ô#:ð àØ ðð ñ�õ Ð<Ð< FÐ<Ñ<Ô<Ñ<Ô<Ð<å(ð 
ð 
ð 
Ø�ð
à"�lð
ð .Ð-ð
ð &Ô2Ð2ð	
ð
 ,˜Oð
ð %Ô0Ð0ð
ð #/Ô"DÐ"Dð
ð +Ô<Ð<ð
ð %<Ð$;ð
ð !4Ô @Ð @ð
ð $:Ð#9ð
ð  3Ô>Ð>ð
ð *FÐ)Eð
ð &9Ô%JÐ%Jð
ð "�zð
ð  "�zð!
ð" "�zð#
ð$ *:Ð)9ð%
ð& )8¨ð'
ð( .Ð-ð)
ð* ,˜Oð+
ð 	
r0   r]   )NNNNNNNNNNNNNNTNN)r5   r6   r7   Ú_tied_weights_keysr   r   r~  ri   r9   r¢   r=  r   r¯  r:   rI   r2   rK   rŸ   r¥   r¦   s   @r.   r'  r'  Ý  s-  ø€ € € € € ð ;Ø0Ø0Ø<ð	ð Ðð!ð !˜{ð !¸B¼IÈÑ<Lð !ð !ð !ð !ð !ð !ðF˜uœ|ð ð ð ð ð
 ð .2Ø48Ø15Ø:>Ø.2Ø.2Ø/3Ø04Ø48Ø8<Ø*.Ø*.Ø*.Ø)-Ø%)Ø#'Ø#'ð%p
ð p
àÔ# dÑ*ðp
ð  Ô*¨TÑ1ðp
ð Ô'¨$Ñ.ð	p
ð
  %Ô0°4Ñ7ðp
ð œ tÑ+ðp
ð œ tÑ+ðp
ð œ¨Ñ,ðp
ð Ô&¨Ñ-ðp
ð $œl¨TÑ1ðp
ð +/°©+ðp
ð ”L 4Ñ'ðp
ð ”L 4Ñ'ðp
ð ”L 4Ñ'ðp
ð   $™;ðp
ð  #ð!p
ð" ˜D‘[ð#p
ð$ ˜D‘[ð%p
ð( 
ˆuŒ|Ô	Ð8Ñ	8ð)p
ð p
ð p
ñ „^ðp
ð p
ð p
ð p
ð p
r0   r'  )r'  rÜ  rb  rY  rW  rL  r€  )Lr8   r¯   ró   r   Údataclassesr   Útypingr   r9   r   Ú r   rT  Úactivationsr	   Úmasking_utilsr
   Úmodeling_layersr   Úmodeling_outputsr   r   Úmodeling_utilsr   Úprocessing_utilsr   Úutilsr   r   r   r   r   r   Úconfiguration_flavar   r   r   r   r   Ú
get_loggerr5   rG  ró  rQ  rR  r  r   r>   rK   r~  r`   ro   r¼   rÞ   r  r  r  r  r"  r-  rC  rL  rb  r€  rW  rY  r±  rÈ  rÑ  rÜ  rþ  rS  r
  r  r'  Ú__all__r<   r0   r.   ú<module>ry     sû  ðð Ð à Ð Ð Ð Ø €€€Ø #Ð #Ð #Ð #Ð #Ð #Ø !Ð !Ð !Ð !Ð !Ð !Ø Ð Ð Ð Ð Ð à €€€Ø Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ -Ð -Ð -Ð -Ð -Ð -Ø &Ð &Ð &Ð &Ð &Ð &Ø jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jÐ jðð ð ð ð ð ð ð ð ð ð ð ð ð ð 
ˆÔ	˜HÑ	%Ô	%€à>Ð àÐ ØÐ à&Ð)9Ñ9Ð<QÑQÐ ð €ððñ ô ð ð
ð 
ð 
ð 
ð 
�{ñ 
ô 
ñ „ñô ð
ð< €ððñ ô ð
 ðð ð ð ð �+ñ ô ñ „ñô ððD €ððñ ô ð ðWtð Wtð Wtð Wtð Wt ñ Wtô Wtñ „ñô ðWtðx_ð _ð _ð _ð _˜2œ9ñ _ô _ð _ðH!ð !ð !ð !ð !�b”iñ !ô !ð !ðH3ð 3ð 3ð 3ð 3˜"œ)ñ 3ô 3ð 3ðl6ð 6ð 6ð 6ð 6˜œñ 6ô 6ð 6ðrð ð ð ð �b”iñ ô ð ð$ð ð ð ð �R”Yñ ô ð ð,ð ð ð ð ˜œ	ñ ô ð ð"ð ð ð ð �"”)ñ ô ð ð )ð )ð )ð )ð )Ð+ñ )ô )ð )ðX$
ð $
ð $
ð $
ð $
�2”9ñ $
ô $
ð $
ðNð ð ð ð �"”)ñ ô ð ð ðSð Sð Sð Sð S˜?ñ Sô Sñ „ðSð6 ðN
ð N
ð N
ð N
ð N
Ð*ñ N
ô N
ñ „ðN
ðb ð\
ð \
ð \
ð \
ð \
Ð)ñ \
ô \
ñ „ð\
ð~ ðL
ð L
ð L
ð L
ð L
Ð/ñ L
ô L
ñ „ðL
ð^ ð^
ð ^
ð ^
ð ^
ð ^
Ð%ñ ^
ô ^
ñ „ð^
ðB	ð ð ð ð  ¤	ñ ô ð ð*Cð Cð Cð Cð C˜bœiñ Cô Cð Cð"ð ð ð ð  2¤9ñ ô ð ð( €ððñ ô ðw)ð w)ð w)ð w)ð w)Ð-ñ w)ô w)ñô ðw)ðtð ð ð ð  2¤9ñ ô ð ð"ð ð ð ð  ¤	ñ ô ð ð 
ð 
ð 
ð 
ð 
�2”9ñ 
ô 
ð 
ð%9ð %9ð %9ð %9ð %9 ¤ñ %9ô %9ð %9ðP €ððñ ô ð
b
ð b
ð b
ð b
ð b
Ð.ñ b
ô b
ñô ð
b
ðJð ð €€€r0   