Ë
    T^(huþ  ã                   ó€  — d Z ddlZddlZddlmZ ddlmZmZm	Z	m
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mZmZmZmZmZmZ dd
lmZ ddlmZmZmZ ddlm Z m!Z!m"Z"m#Z# ddl$m%Z%  e"jL                  e'«      Z(dZ)dZ*e G d„ de«      «       Z+ G d„ dejX                  «      Z- G d„ dejX                  «      Z. G d„ dejX                  «      Z/ G d„ dejX                  «      Z0 G d„ dejX                  «      Z1 G d„ dejX                  «      Z2 G d„ d ejX                  «      Z3 G d!„ d"ejX                  «      Z4 G d#„ d$ejX                  «      Z5 G d%„ d&ejX                  «      Z6 G d'„ d(e«      Z7d)Z8d*Z9d+Z: e d,e8«       G d-„ d.e7«      «       Z; G d/„ d0ejX                  «      Z< e d1e8«       G d2„ d3e7«      «       Z= G d4„ d5ejX                  «      Z> G d6„ d7ejX                  «      Z? e d8e8«       G d9„ d:e7«      «       Z@ e d;e8«       G d<„ d=e7«      «       ZA e d>e:«       G d?„ d@e7«      «       ZB e dAe8«       G dB„ dCe7«      «       ZCg dD¢ZDy)EzPyTorch ViLT model.é    N)Ú	dataclass)ÚListÚOptionalÚTupleÚUnion)Únn)ÚCrossEntropyLossé   )ÚACT2FN)ÚBaseModelOutputÚBaseModelOutputWithPoolingÚMaskedLMOutputÚModelOutputÚSequenceClassifierOutputÚTokenClassifierOutput)ÚPreTrainedModel)Ú find_pruneable_heads_and_indicesÚmeshgridÚprune_linear_layer)Úadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstringsé   )Ú
ViltConfigr   zdandelin/vilt-b32-mlmc                   óÊ   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eeeej                           ed<   dZeeeej                           ed<   y)Ú(ViltForImagesAndTextClassificationOutputa�  
    Class for outputs of [`ViltForImagesAndTextClassification`].

    Args:
        loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
            Classification (or regression if config.num_labels==1) loss.
        logits (`torch.FloatTensor` of shape `(batch_size, config.num_labels)`):
            Classification (or regression if config.num_labels==1) scores (before SoftMax).
        hidden_states (`List[tuple(torch.FloatTensor)]`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            List of tuples of `torch.FloatTensor` (one for each image-text pair, each tuple containing the output of
            the embeddings + one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
            Hidden-states of the model at the output of each layer plus the initial embedding outputs.
        attentions (`List[tuple(torch.FloatTensor)]`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            List of tuples of `torch.FloatTensor` (one for each image-text pair, each tuple containing the attention
            weights of shape `(batch_size, num_heads, sequence_length, sequence_length)`. Attentions weights after the
            attention softmax, used to compute the weighted average in the self-attention heads.
    NÚlossÚlogitsÚhidden_statesÚ
attentions)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   ÚtorchÚFloatTensorÚ__annotations__r   r    r   r   r!   © ó    úd/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/vilt/modeling_vilt.pyr   r   4   sq   … ñð$ )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø*.€FˆH�U×&Ñ&Ñ'Ó.Ø>B€M�8˜D  u×'8Ñ'8Ñ!9Ñ:Ñ;ÓBØ;?€J�˜˜e E×$5Ñ$5Ñ6Ñ7Ñ8Ô?r*   r   c                   ó4   ‡ — e Zd ZdZˆ fd„Zdd„Z	 dd„Zˆ xZS )ÚViltEmbeddingsz¢
    Construct the text and patch embeddings.

    Text embeddings are equivalent to BERT embeddings.

    Patch embeddings are equivalent to ViT embeddings.
    c                 ó,  •— t         ‰| �  «        t        |«      | _        t	        j
                  t        j                  dd|j                  «      «      | _	        t        |«      | _        | j                  j                  }t	        j
                  t        j                  d|dz   |j                  «      «      | _        t	        j                  |j                  |j                  «      | _        t	        j"                  |j$                  «      | _        || _        y ©Nr   )ÚsuperÚ__init__ÚTextEmbeddingsÚtext_embeddingsr   Ú	Parameterr&   ÚzerosÚhidden_sizeÚ	cls_tokenÚViltPatchEmbeddingsÚpatch_embeddingsÚnum_patchesÚposition_embeddingsÚ	EmbeddingÚmodality_type_vocab_sizeÚtoken_type_embeddingsÚDropoutÚhidden_dropout_probÚdropoutÚconfig)ÚselfrB   r:   Ú	__class__s      €r+   r1   zViltEmbeddings.__init__W   sÄ   ø€ Ü‰ÑÔô  .¨fÓ5ˆÔäŸ™¤e§k¡k°!°Q¸×8JÑ8JÓ&KÓLˆŒÜ 3°FÓ ;ˆÔØ×+Ñ+×7Ñ7ˆÜ#%§<¡<´·±¸A¸{ÈQ¹ÐPV×PbÑPbÓ0cÓ#dˆÔ ä%'§\¡\°&×2QÑ2QÐSY×SeÑSeÓ%fˆÔ"Ü—z‘z &×"<Ñ"<Ó=ˆŒØˆ�r*   c                 óL  — | j                   j                  j                  j                  \  }}}}| j                  |«      }|d d …d d d …d d …f   j	                  «       }t
        j                  j                  ||j                  d   |j                  d   f¬«      j                  «       }|d d …df   j                  d¬«      d d …df   }	|d d …df   j                  d¬«      d d …df   }
|j                  \  }}}}| j                  j                  | j                  j                  z  }| j                  d d …dd …d d …f   j                  dd«      j                  d|||«      }t!        j"                  t%        |	|
«      D ��cg c]R  \  }}t
        j                  j'                  t
        j                  j                  |||fdd¬	«      d||z
  d||z
  f«      ‘ŒT c}}d¬«      }|j)                  d«      j                  dd«      }|j)                  d«      j                  dd«      }t!        j*                  t-        t!        j.                  |j                  d
   «      t!        j.                  |j                  d   «      d¬«      d¬«      j1                  |j2                  ¬«      }|d d d d …d d …d d …f   }|j5                  |j                  d   |j                  d   ddd«      }|j)                  dd«      }|j)                  d«      }|dk  s|�t7        |t8        «      s|	|
z  }|j;                  «       }n|	|
z  }t=        |j;                  «       |«      }|j?                  d¬«      }d|z
  j?                  d¬«      }|d d …df   jA                  «       }|D �cg c]  }||d d …df   |k(     ‘Œ }}|D �cg c]  }||d d …df   |k(     ‘Œ }}|D �cg c]  }|jC                  d«      ‘Œ }}|D �cg c]  }|jC                  d«      ‘Œ }}|D �cg c]  }||z
  ‘Œ	 }}g } tE        t%        |||«      «      D ]Ç  \  }!\  }}"}#|#dk  rOt!        jF                  t!        jH                  |«      j	                  «       |«      }$| jK                  ||!   |$   «       Œ^t!        jF                  t!        jH                  |"«      j	                  «       |#d¬«      }%| jK                  t!        j"                  ||!   ||!   |%   gd¬«      «       ŒÉ t!        j"                  | d¬«      } || d d …df   | d d …df   f   j                  |d|«      }|| d d …df   | d d …df   f   j                  |d«      }|| d d …df   | d d …df   f   j                  |dd«      }|| d d …df   | d d …df   f   j                  |d|«      }| jL                  j5                  |dd«      }&t!        j"                  |&|fd¬«      }t!        j"                  | j                  d d …dd d …f   d d …d d d …f   j5                  |dd«      |fd¬«      }||z   }| jO                  |«      }t!        j"                  t!        jH                  |j                  d   d«      j1                  |«      |gd¬«      }|||||fffS c c}}w c c}w c c}w c c}w c c}w c c}w )Né   r
   )Úsizer   r   ©ÚdimÚbilinearT)rG   ÚmodeÚalign_cornerséþÿÿÿéÿÿÿÿÚij)Úindexing©ÚdeviceF)Úas_tuple)Úreplacement)(r9   Ú
projectionÚweightÚshapeÚfloatr   Ú
functionalÚinterpolateÚlongÚsumrB   Ú
image_sizeÚ
patch_sizer;   Ú	transposeÚviewr&   ÚcatÚzipÚpadÚflattenÚstackr   ÚarangeÚtorR   ÚexpandÚ
isinstanceÚintÚmaxÚminÚnonzeroÚuniquerG   Ú	enumerateÚmultinomialÚonesÚappendr7   rA   )'rC   Úpixel_valuesÚ
pixel_maskÚmax_image_lengthÚ_ÚphÚpwÚxÚx_maskÚx_hÚx_wÚ
batch_sizeÚnum_channelsÚheightÚwidthÚ	patch_dimÚspatial_posÚhÚwÚ	pos_embedÚpatch_indexÚeffective_resolutionÚ	valid_idxÚnon_valid_idxÚunique_rowsÚuÚvalid_row_idxÚnon_valid_row_idxÚvÚ
valid_numsÚnon_valid_numsÚpad_numsÚselectÚiÚnvÚpÚvalid_choiceÚ
pad_choiceÚ
cls_tokenss'                                          r+   Úvisual_embedzViltEmbeddings.visual_embedf   sR  € Ø×,Ñ,×7Ñ7×>Ñ>×DÑD‰ˆˆ1ˆb�"à×!Ñ! ,Ó/ˆØšA˜t¢Qª˜MÑ*×0Ñ0Ó2ˆÜ—‘×*Ñ*¨6¸¿¹À¹ÀQÇWÁWÈQÁZÐ8PÐ*ÓQ×VÑVÓXˆØ’Q˜�T‰l×Ñ 1ÐÓ%¢a¨ dÑ+ˆØ’Q˜�T‰l×Ñ 1ÐÓ%¢a¨ dÑ+ˆà23·'±'Ñ/ˆ
�L &¨%Ø—K‘K×*Ñ*¨d¯k©k×.DÑ.DÑDˆ	Ø×.Ñ.ªq°!±"²a¨xÑ8×BÑBÀ1ÀaÓH×MÑMÈaÐQ]Ð_hÐjsÓtˆÜ—I‘Iô    S›M÷ñ �A�qô —‘×!Ñ!Ü—M‘M×-Ñ-Ø#Ø ˜VØ'Ø&*ð	 .ó ð ˜ ™	 1 f¨q¡jÐ1õóð ô
ˆ	ð  ×%Ñ% aÓ(×2Ñ2°1°aÓ8ˆ	Ø�I‰I�a‹L×"Ñ" 1 aÓ(ˆä—k‘kÜ”U—\‘\ &§,¡,¨rÑ"2Ó3´U·\±\À&Ç,Á,ÈrÑBRÓ5SÐ^bÔcÐikô
ç
‰"�F—M‘Mˆ"Ó
"ð 	ð " $¨ªa²²AÐ"5Ñ6ˆØ!×(Ñ(¨¯©°a©¸&¿,¹,Àq¹/È2ÈrÐSUÓVˆØ!×)Ñ)¨!¨QÓ/ˆØ—‘ Ó"ˆà˜aÒÐ#3Ð#;Ä:ÐN^Ô`cÔCdð
 $'¨¡9Ð Ø3×7Ñ7Ó9Ñà#&¨¡9Ð Ü"Ð#7×#;Ñ#;Ó#=Ð?OÓPÐà—N‘N¨E�NÓ2ˆ	Ø˜V™×,Ñ,°eÐ,Ó<ˆØ¢ 1 ‘o×,Ñ,Ó.ˆØBMÖN¸Q˜ 9ªQ°¨T¡?°aÑ#7Ó8ÐNˆÐNØNYÖZÈ˜]¨=º¸A¸Ñ+>À!Ñ+CÓDÐZÐÐZà)6Ö7 A�a—f‘f˜Q•iÐ7ˆ
Ð7Ø->Ö?¨˜!Ÿ&™& �)Ð?ˆÐ?Ø2<Ö=¨QÐ$ qÓ(Ð=ˆÐ=àˆÜ&¤s¨:°~ÀxÓ'PÓQò 	f‰MˆA‰z��2�qØ�AŠvÜ$×0Ñ0´·±¸A³×1DÑ1DÓ1FÐHXÓY�Ø—‘˜m¨AÑ.¨|Ñ<Õ=ä"×.Ñ.¬u¯z©z¸"«~×/CÑ/CÓ/EÀqÐVZÔ[�
Ø—‘œeŸi™i¨°qÑ)9Ð;LÈQÑ;OÐPZÑ;[Ð(\ÐbcÔdÕeð	fô —‘˜6 qÔ)ˆØˆf’Q˜�T‰l˜F¢1 a 4™LÐ(Ñ)×.Ñ.¨z¸2¸|ÓLˆØ˜šq !˜t™ fªQ°¨T¡lÐ2Ñ3×8Ñ8¸ÀRÓHˆà! &ª¨A¨¡,°²q¸!°t±Ð"<Ñ=×BÑBÀ:ÈrÐSTÓUˆØ˜f¢Q¨ T™l¨F²1°a°4©LÐ8Ñ9×>Ñ>¸zÈ2È|Ó\ˆ	à—^‘^×*Ñ*¨:°r¸2Ó>ˆ
Ü�I‰I�z 1�o¨1Ô-ˆÜ—I‘IØ×%Ñ%¢a¨ªA gÑ.ªq°$º¨zÑ:×AÑAÀ*ÈbÐRTÓUÐW`ÐaÐghô
ˆ	ð �	‰MˆØ�L‰L˜‹Oˆä—‘œEŸJ™J v§|¡|°A¡¸Ó:×=Ñ=¸fÓEÀvÐNÐTUÔVˆà�&˜;¨°¨Ð8Ð8Ð8ùóSùòP OùÚZùâ7ùÚ?ùÚ=s%   Å?AZ
ÎZÎ+ZÏZÏ%ZÐZ!c	           	      ó(  — | j                  |||¬«      }	|€-| j                  ||| j                  j                  ¬«      \  }}
}n|j	                  d«      }
|€d}|	| j                  t        j                  |t        j                  |	j                  ¬«      «      z   }	|| j                  t        j                  |
|t        j                  |	j                  ¬«      «      z   }t        j                  |	|gd¬«      }t        j                  ||
gd¬«      }||fS )N)Ú	input_idsÚtoken_type_idsÚinputs_embeds)ru   r   ©ÚdtyperR   rH   )r3   r™   rB   ru   rd   r>   r&   Ú
zeros_liker[   rR   Ú	full_likera   )rC   r›   Úattention_maskrœ   rs   rt   r�   Úimage_embedsÚimage_token_type_idxÚtext_embedsÚimage_masksr†   Ú
embeddingsÚmaskss                 r+   ÚforwardzViltEmbeddings.forward¾   s  € ð ×*Ñ*Ø°Èmð +ó 
ˆð
 ÐØ59×5FÑ5FØ˜j¸4¿;¹;×;WÑ;Wð 6Gó 6Ñ2ˆL˜+¡{ð %×,Ñ,¨QÓ/ˆKð  Ð'Ø#$Ð Ø! D×$>Ñ$>Ü×Ñ˜^´5·:±:Àk×FXÑFXÔYó%
ñ 
ˆð $ d×&@Ñ&@Ü�O‰O˜KÐ)=ÄUÇZÁZÐXc×XjÑXjÔkó'
ñ 
ˆô
 —Y‘Y ¨\Ð:ÀÔBˆ
Ü—	‘	˜>¨;Ð7¸QÔ?ˆà˜5Ð Ð r*   )éÈ   )r   )r"   r#   r$   r%   r1   r™   r©   Ú__classcell__©rD   s   @r+   r-   r-   N   s   ø„ ñôóV9ðB ÷'!r*   r-   c                   ó*   ‡ — e Zd ZdZˆ fd„Zdd„Zˆ xZS )r2   zGConstruct 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                  «      | _        t#        |dd«      | _        | j'                  dt)        j*                  |j                  «      j-                  d«      d¬«       | j'                  d	t)        j.                  | j0                  j3                  «       t(        j4                  ¬
«      d¬«       y )N)Úpadding_idx©ÚepsÚposition_embedding_typeÚabsoluteÚposition_ids)r   rN   F)Ú
persistentrœ   ©rŸ   )r0   r1   r   r<   Ú
vocab_sizer6   Úpad_token_idÚword_embeddingsÚmax_position_embeddingsr;   Útype_vocab_sizer>   Ú	LayerNormÚlayer_norm_epsr?   r@   rA   Úgetattrr²   Úregister_bufferr&   rf   rh   r5   r´   rG   r[   ©rC   rB   rD   s     €r+   r1   zTextEmbeddings.__init__ë   s/  ø€ Ü‰ÑÔÜ!Ÿ|™|¨F×,=Ñ,=¸v×?QÑ?QÐ_e×_rÑ_rÔsˆÔÜ#%§<¡<°×0NÑ0NÐPV×PbÑPbÓ#cˆÔ Ü%'§\¡\°&×2HÑ2HÈ&×J\ÑJ\Ó%]ˆÔ"ô Ÿ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÜ—z‘z &×"<Ñ"<Ó=ˆŒä'.¨vÐ7PÐR\Ó']ˆÔ$Ø×ÑØœEŸL™L¨×)GÑ)GÓH×OÑOÐPWÓXÐejð 	ô 	
ð 	×ÑØœeŸk™k¨$×*;Ñ*;×*@Ñ*@Ó*BÌ%Ï*É*ÔUÐbgð 	õ 	
r*   c                 óT  — |�|j                  «       }n|j                  «       d d }|d   }|€| j                  d d …d |…f   }|€st        | d«      r-| j                  d d …d |…f   }|j	                  |d   |«      }|}n:t        j                  |t
        j                  | j                  j                  ¬«      }|€| j                  |«      }| j                  |«      }	||	z   }
| j                  dk(  r| j                  |«      }|
|z  }
| j                  |
«      }
| j                  |
«      }
|
S )NrN   r   rœ   r   rž   r³   )rG   r´   Úhasattrrœ   rh   r&   r5   r[   rR   r¹   r>   r²   r;   r¼   rA   )rC   r›   rœ   r´   r�   Úinput_shapeÚ
seq_lengthÚbuffered_token_type_idsÚ buffered_token_type_ids_expandedr>   r§   r;   s               r+   r©   zTextEmbeddings.forwardþ   s=  € ØÐ Ø#Ÿ.™.Ó*‰Kà'×,Ñ,Ó.¨s°Ð3ˆKà  ‘^ˆ
àÐØ×,Ñ,ªQ°°°¨^Ñ<ˆLð
 Ð!Ü�tÐ-Ô.Ø*.×*=Ñ*=ºaÀÀ*À¸nÑ*MÐ'Ø3J×3QÑ3QÐR]Ð^_ÑR`ÐblÓ3mÐ0Ø!A‘ä!&§¡¨[ÄÇ
Á
ÐSW×SdÑSd×SkÑSkÔ!l�àÐ Ø ×0Ñ0°Ó;ˆMØ $× :Ñ :¸>Ó JÐà"Ð%:Ñ:ˆ
Ø×'Ñ'¨:Ò5Ø"&×":Ñ":¸<Ó"HÐØÐ-Ñ-ˆJØ—^‘^ JÓ/ˆ
Ø—\‘\ *Ó-ˆ
ØÐr*   )NNNN©r"   r#   r$   r%   r1   r©   r«   r¬   s   @r+   r2   r2   è   s   ø„ ÙQô
÷& r*   r2   c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )r8   z#
    Image to Patch Embedding.
    c                 óÌ  •— t         ‰| �  «        |j                  |j                  }}|j                  |j
                  }}t        |t        j                  j                  «      r|n||f}t        |t        j                  j                  «      r|n||f}|d   |d   z  |d   |d   z  z  }|| _        || _        || _        || _
        t        j                  ||||¬«      | _        y )Nr   r   )Úkernel_sizeÚstride)r0   r1   r]   r^   r~   r6   ri   ÚcollectionsÚabcÚIterabler:   r   ÚConv2drU   )rC   rB   r]   r^   r~   r6   r:   rD   s          €r+   r1   zViltPatchEmbeddings.__init__&  sÔ   ø€ Ü‰ÑÔØ!'×!2Ñ!2°F×4EÑ4E�Jˆ
Ø$*×$7Ñ$7¸×9KÑ9K�kˆä#-¨j¼+¿/¹/×:RÑ:RÔ#S‘ZÐZdÐfpÐYqˆ
Ü#-¨j¼+¿/¹/×:RÑ:RÔ#S‘ZÐZdÐfpÐYqˆ
Ø! !‘}¨
°1©Ñ5¸*ÀQ¹-È:ÐVWÉ=Ñ:XÑYˆØ$ˆŒØ$ˆŒØ(ˆÔØ&ˆÔäŸ)™) L°+È:Ð^hÔiˆ�r*   c                 óÞ   — |j                   \  }}}}|| j                  k7  rt        d«      ‚| j                  j                  j
                  }| j                  |j                  |¬«      «      }|S )NzeMake sure that the channel dimension of the pixel values match with the one set in the configuration.r¶   )rW   r~   Ú
ValueErrorrU   rV   rŸ   rg   )rC   rs   r}   r~   r   r€   Útarget_dtypery   s           r+   r©   zViltPatchEmbeddings.forward5  si   € Ø2>×2DÑ2DÑ/ˆ
�L &¨%Ø˜4×,Ñ,Ò,ÜØwóð ð —‘×-Ñ-×3Ñ3ˆØ�O‰O˜LŸO™O°,˜OÓ?Ó@ˆØˆr*   rÇ   r¬   s   @r+   r8   r8   !  s   ø„ ñôjör*   r8   c                   ó,   ‡ — e Zd Zˆ fd„Zd„ Zdd„Zˆ xZS )ÚViltSelfAttentionc                 ó  •— t         ‰| �  «        |j                  |j                  z  dk7  r2t	        |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                   «      | _        y )Nr   Úembedding_sizezThe hidden size z4 is not a multiple of the number of attention heads ú.©Úbias)r0   r1   r6   Únum_attention_headsrÂ   rÑ   rj   Úattention_head_sizeÚall_head_sizer   ÚLinearÚqkv_biasÚqueryÚkeyÚvaluer?   Úattention_probs_dropout_probrA   rÀ   s     €r+   r1   zViltSelfAttention.__init__A  s.  ø€ Ü‰ÑÔØ×Ñ × :Ñ :Ñ:¸aÒ?ÌÐPVÐXhÔHiÜØ" 6×#5Ñ#5Ð"6ð 7Ø×3Ñ3Ð4°Að7óð ð
 $*×#=Ñ#=ˆÔ Ü#& v×'9Ñ'9¸F×<VÑ<VÑ'VÓ#WˆÔ Ø!×5Ñ5¸×8PÑ8PÑPˆÔä—Y‘Y˜v×1Ñ1°4×3EÑ3EÈFÏOÉOÔ\ˆŒ
Ü—9‘9˜V×/Ñ/°×1CÑ1CÈ&Ï/É/ÔZˆŒÜ—Y‘Y˜v×1Ñ1°4×3EÑ3EÈFÏOÉOÔ\ˆŒ
ä—z‘z &×"EÑ"EÓFˆ�r*   c                 ó    — |j                  «       d d | j                  | j                  fz   } |j                  |Ž }|j	                  dddd«      S )NrN   r   rF   r   r
   )rG   rÚ   rÛ   r`   Úpermute)rC   ry   Únew_x_shapes      r+   Útranspose_for_scoresz&ViltSelfAttention.transpose_for_scoresS  sN   € Ø—f‘f“h˜s �m t×'?Ñ'?À×AYÑAYÐ&ZÑZˆØˆA�F‰F�KÐ ˆØ�y‰y˜˜A˜q !Ó$Ð$r*   c                 ó¶  — | j                  |«      }| j                  | j                  |«      «      }| j                  | j                  |«      «      }| j                  |«      }t	        j
                  ||j                  dd«      «      }	|	t        j                  | j                  «      z  }	|�|	|z   }	 t        j                  d¬«      |	«      }
| j                  |
«      }
|�|
|z  }
t	        j
                  |
|«      }|j                  dddd«      j                  «       }|j                  «       d d | j                   fz   } |j"                  |Ž }|r||
f}|S |f}|S )NrN   rM   rH   r   rF   r   r
   )rß   ræ   rà   rá   r&   Úmatmulr_   ÚmathÚsqrtrÛ   r   ÚSoftmaxrA   rä   Ú
contiguousrG   rÜ   r`   )rC   r    r¢   Ú	head_maskÚoutput_attentionsÚmixed_query_layerÚ	key_layerÚvalue_layerÚquery_layerÚattention_scoresÚattention_probsÚcontext_layerÚnew_context_layer_shapeÚoutputss                 r+   r©   zViltSelfAttention.forwardX  sa  € Ø ŸJ™J }Ó5Ðà×-Ñ-¨d¯h©h°}Ó.EÓFˆ	Ø×/Ñ/°·
±
¸=Ó0IÓJˆØ×/Ñ/Ð0AÓBˆô !Ÿ<™<¨°Y×5HÑ5HÈÈRÓ5PÓQÐØ+¬d¯i©i¸×8PÑ8PÓ.QÑQÐØÐ%à/°.Ñ@Ðð -œ"Ÿ*™*¨Ô,Ð-=Ó>ˆð Ÿ,™, Ó7ˆð Ð Ø-°	Ñ9ˆOäŸ™ _°kÓBˆà%×-Ñ-¨a°°A°qÓ9×DÑDÓFˆØ"/×"4Ñ"4Ó"6°s¸Ð";¸t×?QÑ?QÐ>SÑ"SÐØ*˜×*Ñ*Ð,CÐDˆá6G�= /Ð2ˆàˆð O\ÐM]ˆàˆr*   ©NNF)r"   r#   r$   r1   ræ   r©   r«   r¬   s   @r+   rÔ   rÔ   @  s   ø„ ôGò$%÷
!r*   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 )	ÚViltSelfOutputz¡
    The residual connection is defined in ViltLayer instead of here (as is the case with other models), due to the
    layernorm applied before each block.
    rB   ÚreturnNc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  |j                  «      | _        y ©N)	r0   r1   r   rÝ   r6   Údenser?   r@   rA   rÀ   s     €r+   r1   zViltSelfOutput.__init__ƒ  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r*   r    Úinput_tensorc                 óJ   — | j                  |«      }| j                  |«      }|S rý   ©rþ   rA   ©rC   r    rÿ   s      r+   r©   zViltSelfOutput.forwardˆ  s$   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆàÐr*   )
r"   r#   r$   r%   r   r1   r&   ÚTensorr©   r«   r¬   s   @r+   rú   rú   }  sD   ø„ ñð
>˜zð >¨dõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r*   rú   c                   ó,   ‡ — e Zd Zˆ fd„Zd„ Zdd„Zˆ xZS )ÚViltAttentionc                 ó€   •— t         ‰| �  «        t        |«      | _        t	        |«      | _        t        «       | _        y rý   )r0   r1   rÔ   Ú	attentionrú   ÚoutputÚsetÚpruned_headsrÀ   s     €r+   r1   zViltAttention.__init__�  s0   ø€ Ü‰ÑÔÜ*¨6Ó2ˆŒÜ$ VÓ,ˆŒÜ›EˆÕr*   c                 ó>  — t        |«      dk(  ry t        || j                  j                  | j                  j                  | j
                  «      \  }}t        | j                  j                  |«      | j                  _        t        | j                  j                  |«      | j                  _        t        | j                  j                  |«      | j                  _	        t        | j                  j                  |d¬«      | j                  _        | j                  j                  t        |«      z
  | j                  _        | j                  j                  | j                  j                  z  | j                  _        | j
                  j                  |«      | _        y )Nr   r   rH   )Úlenr   r  rÚ   rÛ   r
  r   rß   rà   rá   r  rþ   rÜ   Úunion)rC   ÚheadsÚindexs      r+   Úprune_headszViltAttention.prune_heads–  s  € Üˆu‹:˜Š?ØÜ7Ø�4—>‘>×5Ñ5°t·~±~×7YÑ7YÐ[_×[lÑ[ló
‰ˆˆuô
  2°$·.±.×2FÑ2FÈÓNˆ�‰ÔÜ/°·±×0BÑ0BÀEÓJˆ�‰ÔÜ1°$·.±.×2FÑ2FÈÓNˆ�‰ÔÜ.¨t¯{©{×/@Ñ/@À%ÈQÔOˆ�‰Ôð .2¯^©^×-OÑ-OÔRUÐV[ÓR\Ñ-\ˆ�‰Ô*Ø'+§~¡~×'IÑ'IÈDÏNÉN×LnÑLnÑ'nˆ�‰Ô$Ø ×-Ñ-×3Ñ3°EÓ:ˆÕr*   c                 ój   — | j                  ||||«      }| j                  |d   |«      }|f|dd  z   }|S )Nr   r   )r  r  )rC   r    r¢   rí   rî   Úself_outputsÚattention_outputr÷   s           r+   r©   zViltAttention.forward¨  sE   € Ø—~‘~ m°^ÀYÐPaÓbˆàŸ;™; |°A¡¸ÓFÐà#Ð%¨°Q°RÐ(8Ñ8ˆØˆr*   rø   )r"   r#   r$   r1   r  r©   r«   r¬   s   @r+   r  r  �  s   ø„ ô"ò;÷$r*   r  c                   ó`   ‡ — e Zd Zdeddfˆ fd„Zdej                  dej                  fd„Zˆ xZS )ÚViltIntermediaterB   rû   Nc                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y rý   )r0   r1   r   rÝ   r6   Úintermediate_sizerþ   ri   Ú
hidden_actÚstrr   Úintermediate_act_fnrÀ   s     €r+   r1   zViltIntermediate.__init__³  s]   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3KÑ3KÓLˆŒ
Ü�f×'Ñ'¬Ô-Ü'-¨f×.?Ñ.?Ñ'@ˆDÕ$à'-×'8Ñ'8ˆDÕ$r*   r    c                 óJ   — | j                  |«      }| j                  |«      }|S rý   )rþ   r  ©rC   r    s     r+   r©   zViltIntermediate.forward»  s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆàÐr*   ©	r"   r#   r$   r   r1   r&   r  r©   r«   r¬   s   @r+   r  r  ²  s1   ø„ ð9˜zð 9¨dõ 9ð U§\¡\ð °e·l±l÷ r*   r  c                   óx   ‡ — 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 )Ú
ViltOutputrB   rû   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j                  «      | _	        y rý   )
r0   r1   r   rÝ   r  r6   rþ   r?   r@   rA   rÀ   s     €r+   r1   zViltOutput.__init__Ä  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×7Ñ7¸×9KÑ9KÓLˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r*   r    rÿ   c                 óT   — | j                  |«      }| j                  |«      }||z   }|S rý   r  r  s      r+   r©   zViltOutput.forwardÉ  s.   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆà%¨Ñ4ˆàÐr*   r  r¬   s   @r+   r  r  Ã  s?   ø„ ð>˜zð >¨dõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r*   r  c                   ó*   ‡ — e Zd ZdZˆ fd„Zdd„Zˆ xZS )Ú	ViltLayerz?This corresponds to the Block class in the timm implementation.c                 ór  •— t         ‰| �  «        |j                  | _        d| _        t	        |«      | _        t        |«      | _        t        |«      | _	        t        j                  |j                  |j                  ¬«      | _        t        j                  |j                  |j                  ¬«      | _        y )Nr   r°   )r0   r1   Úchunk_size_feed_forwardÚseq_len_dimr  r  r  Úintermediater  r  r   r¼   r6   r½   Úlayernorm_beforeÚlayernorm_afterrÀ   s     €r+   r1   zViltLayer.__init__Õ  s‡   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ& vÓ.ˆŒÜ,¨VÓ4ˆÔÜ  Ó(ˆŒÜ "§¡¨V×-?Ñ-?ÀV×EZÑEZÔ [ˆÔÜ!Ÿ|™|¨F×,>Ñ,>ÀF×DYÑDYÔZˆÕr*   c                 ó  — | j                  | j                  |«      |||¬«      }|d   }|dd  }||j                  |j                  «      z   }| j	                  |«      }| j                  |«      }| j                  ||«      }|f|z   }|S )N)rî   r   r   )r  r(  rg   rR   r)  r'  r  )	rC   r    r¢   rí   rî   Úself_attention_outputsr  r÷   Úlayer_outputs	            r+   r©   zViltLayer.forwardß  s©   € Ø!%§¡Ø×!Ñ! -Ó0ØØØ/ð	 "0ó "
Ðð 2°!Ñ4ÐØ(¨¨Ð,ˆð )¨=×+;Ñ+;Ð<L×<SÑ<SÓ+TÑTˆð ×+Ñ+¨MÓ:ˆØ×(Ñ(¨Ó6ˆð —{‘{ <°Ó?ˆà�/ GÑ+ˆàˆr*   rø   rÇ   r¬   s   @r+   r#  r#  Ò  s   ø„ ÙIô[÷r*   r#  c                   ó0   ‡ — e Zd Zˆ fd„Z	 	 	 	 	 dd„Zˆ xZS )ÚViltEncoderc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w )NF)
r0   r1   rB   r   Ú
ModuleListÚrangeÚnum_hidden_layersr#  ÚlayerÚgradient_checkpointing)rC   rB   rv   rD   s      €r+   r1   zViltEncoder.__init__ù  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]¼uÀV×E]ÑE]Ó?^Ö#_¸!¤I¨fÕ$5Ò#_Ó`ˆŒ
Ø&+ˆÕ#ùò $`s   ½A#c                 óx  — |rdnd }|rdnd }t        | j                  «      D ]j  \  }	}
|r||fz   }|�||	   nd }| j                  r,| j                  r | j	                  |
j
                  ||||«      }n |
||||«      }|d   }|sŒb||d   fz   }Œl |r||fz   }|st        d„ |||fD «       «      S t        |||¬«      S )Nr)   r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrý   r)   )Ú.0rŽ   s     r+   ú	<genexpr>z&ViltEncoder.forward.<locals>.<genexpr>%  s   è ø€ Òm˜qÐ_`Ñ_lœÑmùs   ‚Š)Úlast_hidden_stater    r!   )ro   r3  r4  ÚtrainingÚ_gradient_checkpointing_funcÚ__call__Útupler   )rC   r    r¢   rí   rî   Úoutput_hidden_statesÚreturn_dictÚall_hidden_statesÚall_self_attentionsr“   Úlayer_moduleÚlayer_head_maskÚlayer_outputss                r+   r©   zViltEncoder.forwardÿ  s  € ñ #7™B¸DÐÙ$5™b¸4Ðä(¨¯©Ó4ò 	P‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à.7Ð.C˜i¨šlÈˆOà×*Ò*¨t¯}ª}Ø $× AÑ AØ ×)Ñ)Ø!Ø"Ø#Ø%ó!‘ñ !-¨]¸NÈOÐ]nÓ o�à)¨!Ñ,ˆMâ Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð)	Pñ,  Ø 1°]Ð4DÑ DÐáÜÑm ]Ð4EÐGZÐ$[ÔmÓmÐmÜØ+Ø+Ø*ô
ð 	
r*   )NNFFT©r"   r#   r$   r1   r©   r«   r¬   s   @r+   r.  r.  ø  s   ø„ ô,ð ØØØ"Ø÷+
r*   r.  c                   ó*   — e Zd ZdZeZdZdZddgZd„ Z	y)ÚViltPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚviltTr-   rÔ   c                 ó"  — t        |t        j                  t        j                  f«      rm|j                  j
                  j                  d| j                  j                  ¬«       |j                  �%|j                  j
                  j                  «        yyt        |t        j                  «      rz|j                  j
                  j                  d| j                  j                  ¬«       |j                  �2|j                  j
                  |j                     j                  «        yyt        |t        j                  «      rJ|j                  j
                  j                  «        |j                  j
                  j                  d«       yy)zInitialize the weightsg        )ÚmeanÚstdNg      ð?)ri   r   rÝ   rÏ   rV   ÚdataÚnormal_rB   Úinitializer_rangerÙ   Úzero_r<   r¯   r¼   Úfill_)rC   Úmodules     r+   Ú_init_weightsz!ViltPreTrainedModel._init_weights8  s  € ä�fœrŸy™y¬"¯)©)Ð4Ô5ð �M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡Ô-Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ×!Ñ!Ð-Ø—‘×"Ñ" 6×#5Ñ#5Ñ6×<Ñ<Õ>ð .ä˜¤§¡Ô-Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)ð .r*   N)
r"   r#   r$   r%   r   Úconfig_classÚbase_model_prefixÚsupports_gradient_checkpointingÚ_no_split_modulesrR  r)   r*   r+   rG  rG  -  s+   „ ñð
 €LØÐØ&*Ð#Ø)Ð+>Ð?Ðó*r*   rG  aH  
    This model is a PyTorch `torch.nn.Module <https://pytorch.org/docs/stable/nn.html#torch.nn.Module>`_ subclass. Use
    it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage and
    behavior.

    Parameters:
        config ([`ViltConfig`]): Model configuration class with all the parameters of the model.
            Initializing with a config file does not load the weights associated with the model, only the
            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
aÍ  
    Args:
        input_ids (`torch.LongTensor` of shape `({0})`):
            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)

        attention_mask (`torch.FloatTensor` of shape `({0})`, *optional*):
            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
            - 1 for tokens that are **not masked**,
            - 0 for tokens that are **masked**.
            [What are attention masks?](../glossary#attention-mask)

        token_type_ids (`torch.LongTensor` of shape `({0})`, *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)

        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See
            [`ViltImageProcessor.__call__`] for details.

        pixel_mask (`torch.LongTensor` of shape `(batch_size, height, width)`, *optional*):
            Mask to avoid performing attention on padding pixel values. Mask values selected in `[0, 1]`:

            - 1 for pixels that are real (i.e. **not masked**),
            - 0 for pixels that are padding (i.e. **masked**).
            `What are attention masks? <../glossary.html#attention-mask>`__

        head_mask (`torch.FloatTensor` of shape `(num_heads,)` or `(num_layers, num_heads)`, *optional*):
            Mask to nullify selected heads of the self-attention modules. Mask values selected in `[0, 1]`:
            - 1 indicates the head is **not masked**,
            - 0 indicates the head is **masked**.

        inputs_embeds (`torch.FloatTensor` of shape `({0}, hidden_size)`, *optional*):
            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
            is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
            model's internal embedding lookup matrix.

        image_embeds (`torch.FloatTensor` of shape `(batch_size, num_patches, hidden_size)`, *optional*):
            Optionally, instead of passing `pixel_values`, you can choose to directly pass an embedded representation.
            This is useful if you want more control over how to convert `pixel_values` into patch embeddings.

        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
añ  
    Args:
        input_ids (`torch.LongTensor` of shape `({0})`):
            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)

        attention_mask (`torch.FloatTensor` of shape `({0})`, *optional*):
            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
            - 1 for tokens that are **not masked**,
            - 0 for tokens that are **masked**.
            [What are attention masks?](../glossary#attention-mask)

        token_type_ids (`torch.LongTensor` of shape `({0})`, *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)

        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_images, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See
            [`ViltImageProcessor.__call__`] for details.

        pixel_mask (`torch.LongTensor` of shape `(batch_size, num_images, height, width)`, *optional*):
            Mask to avoid performing attention on padding pixel values. Mask values selected in `[0, 1]`:

            - 1 for pixels that are real (i.e. **not masked**),
            - 0 for pixels that are padding (i.e. **masked**).
            `What are attention masks? <../glossary.html#attention-mask>`__

        head_mask (`torch.FloatTensor` of shape `(num_heads,)` or `(num_layers, num_heads)`, *optional*):
            Mask to nullify selected heads of the self-attention modules. Mask values selected in `[0, 1]`:
            - 1 indicates the head is **not masked**,
            - 0 indicates the head is **masked**.

        inputs_embeds (`torch.FloatTensor` of shape `({0}, hidden_size)`, *optional*):
            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
            is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
            model's internal embedding lookup matrix.

        image_embeds (`torch.FloatTensor` of shape `(batch_size, num_images, num_patches, hidden_size)`, *optional*):
            Optionally, instead of passing `pixel_values`, you can choose to directly pass an embedded representation.
            This is useful if you want more control over how to convert `pixel_values` into patch embeddings.

        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
z^The bare ViLT Model transformer outputting raw hidden-states without any specific head on top.c                    óÄ  ‡ — e Zd Zdˆ fd„	Zd„ Zd„ Zd„ Z ee«       e	e
e¬«      	 	 	 	 	 	 	 	 	 	 	 	 ddeej                     deej                     deej                     d	eej                     d
eej                     deej                     deej                     deej                     dee   dee   dee   dee   dee
eej                     f   fd„«       «       Zˆ xZS )Ú	ViltModelc                 ó  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _        |rt        |«      nd | _        | j                  «        y ©Nr°   )r0   r1   rB   r-   r§   r.  Úencoderr   r¼   r6   r½   Ú	layernormÚ
ViltPoolerÚpoolerÚ	post_init)rC   rB   Úadd_pooling_layerrD   s      €r+   r1   zViltModel.__init__È  si   ø€ Ü‰Ñ˜Ô ØˆŒä(¨Ó0ˆŒÜ" 6Ó*ˆŒäŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÙ,=”j Ô(À4ˆŒð 	�‰Õr*   c                 óB   — | j                   j                  j                  S rý   ©r§   r3   r¹   ©rC   s    r+   Úget_input_embeddingszViltModel.get_input_embeddingsÕ  s   € Ø�‰×.Ñ.×>Ñ>Ð>r*   c                 ó:   — || j                   j                  _        y rý   rb  )rC   rá   s     r+   Úset_input_embeddingszViltModel.set_input_embeddingsØ  s   € Ø:?ˆ�‰×'Ñ'Õ7r*   c                 ó˜   — |j                  «       D ]7  \  }}| j                  j                  |   j                  j	                  |«       Œ9 y)z�
        Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base
        class PreTrainedModel
        N)Úitemsr[  r3  r  r  )rC   Úheads_to_pruner3  r  s       r+   Ú_prune_headszViltModel._prune_headsÛ  sE   € ð
 +×0Ñ0Ó2ò 	C‰LˆE�5Ø�L‰L×Ñ˜uÑ%×/Ñ/×;Ñ;¸EÕBñ	Cr*   ©Úoutput_typerS  r›   r¢   rœ   rs   rt   rí   r�   r£   r¤   rî   r>  r?  rû   c           
      ó~  — |
�|
n| j                   j                  }
|�|n| j                   j                  }|�|n| j                   j                  }|�|�t	        d«      ‚|�#| j                  ||«       |j                  «       }n!|�|j                  «       dd }nt	        d«      ‚|\  }}|�|j                  n|j                  }|€t        j                  ||f|¬«      }|�|�t	        d«      ‚|€|€t	        d«      ‚|�|j                  d   n|j                  d   }||k7  rt	        d	«      ‚|€Bt        j                  || j                   j                  | j                   j                  f|¬«      }| j                  || j                   j                  «      }| j                  ||||||||	¬
«      \  }}| j                  ||«      }| j!                  ||||
||¬«      }|d   }| j#                  |«      }| j$                  �| j%                  |«      nd}|s
||f|dd z   S t'        |||j(                  |j*                  ¬«      S )aÎ  
        Returns:

        Examples:

        ```python
        >>> from transformers import ViltProcessor, ViltModel
        >>> from PIL import Image
        >>> import requests

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

        >>> processor = ViltProcessor.from_pretrained("dandelin/vilt-b32-mlm")
        >>> model = ViltModel.from_pretrained("dandelin/vilt-b32-mlm")

        >>> inputs = processor(image, text, return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> last_hidden_states = outputs.last_hidden_state
        ```NzDYou cannot specify both input_ids and inputs_embeds at the same timerN   z5You have to specify either input_ids or inputs_embedsrQ   zFYou cannot specify both pixel_values and image_embeds at the same timez7You have to specify either pixel_values or image_embedsr   zAThe text inputs and image inputs need to have the same batch size)r¤   )r¢   rí   rî   r>  r?  r   )r9  Úpooler_outputr    r!   )rB   rî   r>  Úuse_return_dictrÑ   Ú%warn_if_padding_and_no_attention_maskrG   rR   r&   rq   rW   r]   Úget_head_maskr2  r§   Úget_extended_attention_maskr[  r\  r^  r   r    r!   )rC   r›   r¢   rœ   rs   rt   rí   r�   r£   r¤   rî   r>  r?  rÃ   Útext_batch_sizerÄ   rR   Úimage_batch_sizeÚembedding_outputÚextended_attention_maskÚencoder_outputsÚsequence_outputÚpooled_outputs                          r+   r©   zViltModel.forwardã  s‰  € ðN 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ  ]Ð%>ÜÐcÓdÐdØÐ"Ø×6Ñ6°yÀ.ÔQØ#Ÿ.™.Ó*‰KØÐ&Ø'×,Ñ,Ó.¨s°Ð3‰KäÐTÓUÐUà&1Ñ#ˆ˜Ø%.Ð%:�×!Ò!À×@TÑ@TˆàÐ!Ü"ŸZ™Z¨/¸:Ð)FÐPVÔWˆNàÐ#¨Ð(@ÜÐeÓfÐfØÐ! lÐ&:ÜÐVÓWÐWà4@Ð4L˜<×-Ñ-¨aÒ0ÐR^×RdÑRdÐefÑRgÐØ˜Ò.ÜÐ`ÓaÐaØÐÜŸ™Ð%5°t·{±{×7MÑ7MÈtÏ{É{×OeÑOeÐ$fÐouÔvˆJð ×&Ñ& y°$·+±+×2OÑ2OÓPˆ	à+/¯?©?ØØØØØØØØ!5ð ,;ó 	,
Ñ(Ð˜.ð 15×0PÑ0PÐQ_ÐalÓ0mÐàŸ,™,ØØ2ØØ/Ø!5Ø#ð 'ó 
ˆð *¨!Ñ,ˆØŸ.™.¨Ó9ˆØ8<¿¹Ð8O˜Ÿ™ OÔ4ÐUYˆáØ# ]Ð3°oÀaÀbÐ6IÑIÐIä)Ø-Ø'Ø)×7Ñ7Ø&×1Ñ1ô	
ð 	
r*   )T©NNNNNNNNNNNN)r"   r#   r$   r1   rd  rf  rj  r   ÚVILT_INPUTS_DOCSTRINGr   r   Ú_CONFIG_FOR_DOCr   r&   Ú
LongTensorr'   rj   Úboolr   r   r©   r«   r¬   s   @r+   rX  rX  Ã  sy  ø„ õ
ò?ò@òCñ +Ð+@ÓAÙÐ+EÐTcÔdð 15Ø6:Ø59Ø48Ø15Ø15Ø59Ø48Ø.2Ø,0Ø/3Ø&*ñp
à˜E×,Ñ,Ñ-ðp
ð ! ×!2Ñ!2Ñ3ðp
ð ! ×!1Ñ!1Ñ2ð	p
ð
 ˜u×0Ñ0Ñ1ðp
ð ˜U×-Ñ-Ñ.ðp
ð ˜E×-Ñ-Ñ.ðp
ð   × 1Ñ 1Ñ2ðp
ð ˜u×0Ñ0Ñ1ðp
ð ' s™mðp
ð $ D™>ðp
ð ' t™nðp
ð ˜d‘^ðp
ð 
Ð)¨5°×1BÑ1BÑ+CÐCÑ	Dòp
ó eó Bôp
r*   rX  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )r]  c                 ó²   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  «       | _        y rý   )r0   r1   r   rÝ   r6   rþ   ÚTanhÚ
activationrÀ   s     €r+   r1   zViltPooler.__init__Y  s9   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
ÜŸ'™'›)ˆ�r*   c                 ó\   — |d d …df   }| j                  |«      }| j                  |«      }|S )Nr   )rþ   r‚  )rC   r    Úfirst_token_tensorry  s       r+   r©   zViltPooler.forward^  s6   € ð +ª1¨a¨4Ñ0ÐØŸ
™
Ð#5Ó6ˆØŸ™¨Ó6ˆØÐr*   rE  r¬   s   @r+   r]  r]  X  s   ø„ ô$ö
r*   r]  zU
    ViLT Model with a language modeling head on top as done during pretraining.
    c                    óö  ‡ — e Zd ZddgZˆ fd„Zd„ Zd„ Z eej                  d«      «       e
ee¬«      	 	 	 	 	 	 	 	 	 	 	 	 ddeej                     d	eej                      d
eej                     deej                      deej                     deej                      deej                      deej                      deej                     dee   dee   dee   deeeej                      f   fd„«       «       Zˆ xZS )ÚViltForMaskedLMzmlm_score.decoder.weightzmlm_score.decoder.biasc                 ó„   •— t         ‰| �  |«       t        |«      | _        t	        |«      | _        | j                  «        y rý   )r0   r1   rX  rH  ÚViltMLMHeadÚ	mlm_scorer_  rÀ   s     €r+   r1   zViltForMaskedLM.__init__p  s4   ø€ Ü‰Ñ˜Ô ä˜fÓ%ˆŒ	Ü$ VÓ,ˆŒð 	�‰Õr*   c                 ó.   — | j                   j                  S rý   )r‰  Údecoderrc  s    r+   Úget_output_embeddingsz%ViltForMaskedLM.get_output_embeddingsy  s   € Ø�~‰~×%Ñ%Ð%r*   c                 ó\   — || j                   _        |j                  | j                   _        y rý   )r‰  r‹  rÙ   )rC   Únew_embeddingss     r+   Úset_output_embeddingsz%ViltForMaskedLM.set_output_embeddings|  s    € Ø!/ˆ�‰ÔØ,×1Ñ1ˆ�‰Õr*   zbatch_size, sequence_lengthrk  r›   r¢   rœ   rs   rt   rí   r�   r£   Úlabelsrî   r>  r?  rû   c                 óF  — |�|n| j                   j                  }| j                  |||||||||
||¬«      }|dd \  }}|�|j                  d   n|j                  d   }|dd…d|…f   |dd…|d…f   }}| j	                  |«      }d}|	�at        «       }|	j                  |j                  «      }	 ||j                  d| j                   j                  «      |	j                  d«      «      }|s|f|dd z   }|�|f|z   S |S t        |||j                  |j                  ¬«      S )aò	  
        labels (*torch.LongTensor* of shape *(batch_size, sequence_length)*, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in *[-100, 0, ...,
            config.vocab_size]* (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, ..., config.vocab_size]*

        Returns:

        Examples:

        ```python
        >>> from transformers import ViltProcessor, ViltForMaskedLM
        >>> import requests
        >>> from PIL import Image
        >>> import re
        >>> import torch

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)
        >>> text = "a bunch of [MASK] laying on a [MASK]."

        >>> processor = ViltProcessor.from_pretrained("dandelin/vilt-b32-mlm")
        >>> model = ViltForMaskedLM.from_pretrained("dandelin/vilt-b32-mlm")

        >>> # prepare inputs
        >>> encoding = processor(image, text, return_tensors="pt")

        >>> # forward pass
        >>> outputs = model(**encoding)

        >>> tl = len(re.findall("\[MASK\]", text))
        >>> inferred_token = [text]

        >>> # gradually fill in the MASK tokens, one by one
        >>> with torch.no_grad():
        ...     for i in range(tl):
        ...         encoded = processor.tokenizer(inferred_token)
        ...         input_ids = torch.tensor(encoded.input_ids)
        ...         encoded = encoded["input_ids"][0][1:-1]
        ...         outputs = model(input_ids=input_ids, pixel_values=encoding.pixel_values)
        ...         mlm_logits = outputs.logits[0]  # shape (seq_len, vocab_size)
        ...         # only take into account text features (minus CLS and SEP token)
        ...         mlm_logits = mlm_logits[1 : input_ids.shape[1] - 1, :]
        ...         mlm_values, mlm_ids = mlm_logits.softmax(dim=-1).max(dim=-1)
        ...         # only take into account text
        ...         mlm_values[torch.tensor(encoded) != 103] = 0
        ...         select = mlm_values.argmax().item()
        ...         encoded[select] = mlm_ids[select].item()
        ...         inferred_token = [processor.decode(encoded)]

        >>> selected_token = ""
        >>> encoded = processor.tokenizer(inferred_token)
        >>> output = processor.decode(encoded.input_ids[0], skip_special_tokens=True)
        >>> print(output)
        a bunch of cats laying on a couch.
        ```N©
r¢   rœ   rs   rt   rí   r�   r£   rî   r>  r?  rF   r   rN   ©r   r   r    r!   )rB   ro  rH  rW   r‰  r	   rg   rR   r`   r·   r   r    r!   )rC   r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  r÷   rx  ry  Útext_seq_lenÚtext_featuresrv   Ú
mlm_logitsÚmasked_lm_lossÚloss_fctr  s                          r+   r©   zViltForMaskedLM.forward€  s]  € ðR &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø)Ø%Ø!ØØ'Ø%Ø/Ø!5Ø#ð ó 
ˆð *1°°!¨Ñ&ˆ˜à-6Ð-B�y—‘ qÒ)È×H[ÑH[Ð\]ÑH^ˆØ+ªA¨}°¨}Ð,<Ñ=¸ÊqÐR^ÑR_ÐO_Ñ?`�qˆà—^‘^ MÓ2ˆ
àˆØÐÜ'Ó)ˆHà—Y‘Y˜z×0Ñ0Ó1ˆFÙ% j§o¡o°b¸$¿+¹+×:PÑ:PÓ&QÐSY×S^ÑS^Ð_aÓSbÓcˆNáØ �] W¨Q¨R [Ñ0ˆFØ3AÐ3M�^Ð%¨Ñ.ÐYÐSYÐYäØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r*   rz  )r"   r#   r$   Ú_tied_weights_keysr1   rŒ  r�  r   r{  Úformatr   r   r|  r   r&   r}  r'   r~  r   r   r©   r«   r¬   s   @r+   r†  r†  g  s�  ø„ ð 5Ð6NÐOÐôò&ò2ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙ¨>ÈÔXð 15Ø6:Ø59Ø48Ø15Ø15Ø59Ø48Ø-1Ø,0Ø/3Ø&*ñn
à˜E×,Ñ,Ñ-ðn
ð ! ×!2Ñ!2Ñ3ðn
ð ! ×!1Ñ!1Ñ2ð	n
ð
 ˜u×0Ñ0Ñ1ðn
ð ˜U×-Ñ-Ñ.ðn
ð ˜E×-Ñ-Ñ.ðn
ð   × 1Ñ 1Ñ2ðn
ð ˜u×0Ñ0Ñ1ðn
ð ˜×)Ñ)Ñ*ðn
ð $ D™>ðn
ð ' t™nðn
ð ˜d‘^ðn
ð 
ˆ~˜u U×%6Ñ%6Ñ7Ð7Ñ	8òn
ó Yó hôn
r*   r†  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚViltPredictionHeadTransformc                 óh  •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        |j                  t        «      rt        |j                     | _
        n|j                  | _
        t        j                  |j                  |j                  ¬«      | _        y rZ  )r0   r1   r   rÝ   r6   rþ   ri   r  r  r   Útransform_act_fnr¼   r½   rÀ   s     €r+   r1   z$ViltPredictionHeadTransform.__init__ô  s{   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü�f×'Ñ'¬Ô-Ü$*¨6×+<Ñ+<Ñ$=ˆDÕ!à$*×$5Ñ$5ˆDÔ!ÜŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆ�r*   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S rý   )rþ   rž  r¼   r  s     r+   r©   z#ViltPredictionHeadTransform.forwardý  s4   € ØŸ
™
 =Ó1ˆØ×-Ñ-¨mÓ<ˆØŸ™ }Ó5ˆØÐr*   rE  r¬   s   @r+   rœ  rœ  ó  s   ø„ ôUör*   rœ  c                   ó,   ‡ — e Zd Zdˆ fd„	Zd„ Zd„ Zˆ xZS )rˆ  c                 ó|  •— t         ‰| �  «        || _        t        |«      | _        t        j                  |j                  |j                  d¬«      | _	        t        j                  t        j                  |j                  «      «      | _        |�|| j                  _        | j                  | j                  _        y )NFrØ   )r0   r1   rB   rœ  Ú	transformr   rÝ   r6   r·   r‹  r4   r&   r5   rÙ   rV   )rC   rB   rV   rD   s      €r+   r1   zViltMLMHead.__init__  s„   ø€ Ü‰ÑÔØˆŒÜ4°VÓ<ˆŒÜ—y‘y ×!3Ñ!3°V×5FÑ5FÈUÔSˆŒÜ—L‘L¤§¡¨V×->Ñ->Ó!?Ó@ˆŒ	ØÐØ"(ˆD�L‰LÔð !ŸI™Iˆ�‰Õr*   c                 ó:   — | j                   | j                  _         y rý   )rÙ   r‹  rc  s    r+   Ú_tie_weightszViltMLMHead._tie_weights  s   € Ø ŸI™Iˆ�‰Õr*   c                 óJ   — | j                  |«      }| j                  |«      }|S rý   )r¢  r‹  )rC   ry   s     r+   r©   zViltMLMHead.forward  s"   € Ø�N‰N˜1ÓˆØ�L‰L˜‹OˆØˆr*   rý   )r"   r#   r$   r1   r¤  r©   r«   r¬   s   @r+   rˆ  rˆ    s   ø„ õ
&ò&ör*   rˆ  z¶
    Vilt Model transformer with a classifier head on top (a linear layer on top of the final hidden state of the [CLS]
    token) for visual question answering, e.g. for VQAv2.
    c                    óÄ  ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 	 	 	 	 	 	 	 	 	 	 	 dde	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     d	e	e
j                     d
e	e
j                     de	e
j                     de	e   de	e   de	e   deeee
j                     f   fd„«       «       Zˆ xZS )ÚViltForQuestionAnsweringc           	      óÐ  •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        j                  t        j                  |j                  |j                  dz  «      t        j                  |j                  dz  «      t        j                  «       t        j                  |j                  dz  |j                  «      «      | _        | j                  «        y )NrF   )r0   r1   Ú
num_labelsrX  rH  r   Ú
SequentialrÝ   r6   r¼   ÚGELUÚ
classifierr_  rÀ   s     €r+   r1   z!ViltForQuestionAnswering.__init__"  s¥   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ˜fÓ%ˆŒ	ô Ÿ-™-Ü�I‰I�f×(Ñ(¨&×*<Ñ*<¸qÑ*@ÓAÜ�L‰L˜×+Ñ+¨aÑ/Ó0Ü�G‰G‹IÜ�I‰I�f×(Ñ(¨1Ñ,¨f×.?Ñ.?Ó@ó	
ˆŒð 	�‰Õr*   rk  r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  rû   c                 óÄ  — |�|n| j                   j                  }| j                  |||||||||
||¬«      }|r|j                  n|d   }| j	                  |«      }d}|	�K|	j                  |j                  «      }	t        j                  j                  ||	«      |	j                  d   z  }|s|f|dd z   }|�|f|z   S |S t        |||j                  |j                  ¬«      S )a  
        labels (`torch.FloatTensor` of shape `(batch_size, num_labels)`, *optional*):
            Labels for computing the visual question answering loss. This tensor must be either a one-hot encoding of
            all answers that are applicable for a given example in the batch, or a soft encoding indicating which
            answers are applicable, where 1.0 is the highest score.

        Returns:

        Examples:

        ```python
        >>> from transformers import ViltProcessor, ViltForQuestionAnswering
        >>> import requests
        >>> from PIL import Image

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)
        >>> text = "How many cats are there?"

        >>> processor = ViltProcessor.from_pretrained("dandelin/vilt-b32-finetuned-vqa")
        >>> model = ViltForQuestionAnswering.from_pretrained("dandelin/vilt-b32-finetuned-vqa")

        >>> # prepare inputs
        >>> encoding = processor(image, text, return_tensors="pt")

        >>> # forward pass
        >>> outputs = model(**encoding)
        >>> logits = outputs.logits
        >>> idx = logits.argmax(-1).item()
        >>> print("Predicted answer:", model.config.id2label[idx])
        Predicted answer: 2
        ```Nr’  r   rF   r“  )rB   ro  rH  rn  r¬  rg   rR   r   rY   Ú binary_cross_entropy_with_logitsrW   r   r    r!   )rC   r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  r÷   rn  r   r   r  s                     r+   r©   z ViltForQuestionAnswering.forward3  s
  € ðb &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø)Ø%Ø!ØØ'Ø%Ø/Ø!5Ø#ð ó 
ˆñ 2=˜×-Ò-À'È!Á*ˆà—‘ Ó/ˆàˆØÐà—Y‘Y˜vŸ}™}Ó-ˆFÜ—=‘=×AÑAÀ&È&ÓQÐTZ×T`ÑT`ÐabÑTcÑcˆDñ Ø�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä'ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r*   rz  ©r"   r#   r$   r1   r   r{  r   r   r|  r   r&   r}  r'   r~  r   r   r©   r«   r¬   s   @r+   r§  r§    so  ø„ ôñ" +Ð+@ÓAÙÐ+CÐRaÔbð 15Ø6:Ø59Ø48Ø15Ø15Ø59Ø48Ø-1Ø,0Ø/3Ø&*ñS
à˜E×,Ñ,Ñ-ðS
ð ! ×!2Ñ!2Ñ3ðS
ð ! ×!1Ñ!1Ñ2ð	S
ð
 ˜u×0Ñ0Ñ1ðS
ð ˜U×-Ñ-Ñ.ðS
ð ˜E×-Ñ-Ñ.ðS
ð   × 1Ñ 1Ñ2ðS
ð ˜u×0Ñ0Ñ1ðS
ð ˜×)Ñ)Ñ*ðS
ð $ D™>ðS
ð ' t™nðS
ð ˜d‘^ðS
ð 
Ð'¨¨u×/@Ñ/@Ñ)AÐAÑ	BòS
ó có BôS
r*   r§  zË
    Vilt Model transformer with a classifier head on top (a linear layer on top of the final hidden state of the [CLS]
    token) for image-to-text or text-to-image retrieval, e.g. MSCOCO and F30K.
    c                    óÄ  ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 	 	 	 	 	 	 	 	 	 	 	 dde	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     d	e	e
j                     d
e	e
j                     de	e
j                     de	e   de	e   de	e   deeee
j                     f   fd„«       «       Zˆ xZS )ÚViltForImageAndTextRetrievalc                 ó®   •— t         ‰| �  |«       t        |«      | _        t	        j
                  |j                  d«      | _        | j                  «        y r/   )	r0   r1   rX  rH  r   rÝ   r6   Úrank_outputr_  rÀ   s     €r+   r1   z%ViltForImageAndTextRetrieval.__init__“  sC   ø€ Ü‰Ñ˜Ô ä˜fÓ%ˆŒ	ô Ÿ9™9 V×%7Ñ%7¸Ó;ˆÔð 	�‰Õr*   rk  r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  rû   c                 óD  — |�|n| j                   j                  }d}|	�t        d«      ‚| j                  |||||||||
||¬«      }|r|j                  n|d   }| j                  |«      }|s|f|dd z   }|�|f|z   S |S t        |||j                  |j                  ¬«      S )a'  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels are currently not supported.

        Returns:

        Examples:

        ```python
        >>> from transformers import ViltProcessor, ViltForImageAndTextRetrieval
        >>> import requests
        >>> from PIL import Image

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)
        >>> texts = ["An image of two cats chilling on a couch", "A football player scoring a goal"]

        >>> processor = ViltProcessor.from_pretrained("dandelin/vilt-b32-finetuned-coco")
        >>> model = ViltForImageAndTextRetrieval.from_pretrained("dandelin/vilt-b32-finetuned-coco")

        >>> # forward pass
        >>> scores = dict()
        >>> for text in texts:
        ...     # prepare inputs
        ...     encoding = processor(image, text, return_tensors="pt")
        ...     outputs = model(**encoding)
        ...     scores[text] = outputs.logits[0, :].item()
        ```NzTraining is not yet supported.r’  r   rF   r“  )	rB   ro  ÚNotImplementedErrorrH  rn  r³  r   r    r!   )rC   r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  r   r÷   rn  r   r  s                     r+   r©   z$ViltForImageAndTextRetrieval.forwardž  sÜ   € ðZ &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàˆØÐÜ%Ð&FÓGÐGà—)‘)ØØ)Ø)Ø%Ø!ØØ'Ø%Ø/Ø!5Ø#ð ó 
ˆñ 2=˜×-Ò-À'È!Á*ˆà×!Ñ! -Ó0ˆáØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä'ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r*   rz  r¯  r¬   s   @r+   r±  r±  ‹  so  ø„ ô	ñ +Ð+@ÓAÙÐ+CÐRaÔbð 15Ø6:Ø59Ø48Ø15Ø15Ø59Ø48Ø-1Ø,0Ø/3Ø&*ñL
à˜E×,Ñ,Ñ-ðL
ð ! ×!2Ñ!2Ñ3ðL
ð ! ×!1Ñ!1Ñ2ð	L
ð
 ˜u×0Ñ0Ñ1ðL
ð ˜U×-Ñ-Ñ.ðL
ð ˜E×-Ñ-Ñ.ðL
ð   × 1Ñ 1Ñ2ðL
ð ˜u×0Ñ0Ñ1ðL
ð ˜×)Ñ)Ñ*ðL
ð $ D™>ðL
ð ' t™nðL
ð ˜d‘^ðL
ð 
Ð'¨¨u×/@Ñ/@Ñ)AÐAÑ	BòL
ó có BôL
r*   r±  zq
    Vilt Model transformer with a classifier head on top for natural language visual reasoning, e.g. NLVR2.
    c                    óÄ  ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 	 	 	 	 	 	 	 	 	 	 	 dde	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     d	e	e
j                     d
e	e
j                     de	e
j                     de	e   de	e   de	e   deeee
j                     f   fd„«       «       Zˆ xZS )Ú"ViltForImagesAndTextClassificationc           	      óî  •— t         ‰| �  |«       |j                  | _        t        |«      | _        |j
                  }t        j                  t        j                  |j                  |z  |j                  |z  «      t        j                  |j                  |z  «      t        j                  «       t        j                  |j                  |z  |j                  «      «      | _        | j                  «        y rý   )r0   r1   r©  rX  rH  Ú
num_imagesr   rª  rÝ   r6   r¼   r«  r¬  r_  )rC   rB   r¹  rD   s      €r+   r1   z+ViltForImagesAndTextClassification.__init__ö  sµ   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ˜fÓ%ˆŒ	ð ×&Ñ&ˆ
ÜŸ-™-Ü�I‰I�f×(Ñ(¨:Ñ5°v×7IÑ7IÈJÑ7VÓWÜ�L‰L˜×+Ñ+¨jÑ8Ó9Ü�G‰G‹IÜ�I‰I�f×(Ñ(¨:Ñ5°v×7HÑ7HÓIó	
ˆŒð 	�‰Õr*   rk  r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  rû   c                 óª  — |
�|
n| j                   j                  }
|�|n| j                   j                  }|�|n| j                   j                  }|� |j                  dk(  r|j                  d«      }|� |j                  dk(  r|j                  d«      }|�|j                  d   nd}|€|�|j                  d   nd}|| j                   j                  k7  rt        d«      ‚g }|rg nd}|
rg nd}t        |«      D ]·  }| j                  ||||�|dd…|dd…dd…dd…f   nd|�|dd…|dd…dd…f   nd|||�|dd…|dd…dd…f   nd|dz   |
||¬«      }|r|j                  n|d   }|j                  |«       |r|j                  |j                  «       |
sŒ�|j                  |j                  «       Œ¹ t        j                   |d¬«      }| j#                  |«      }d}|	�Wt%        «       }|	j'                  |j(                  «      }	 ||j+                  d| j,                  «      |	j+                  d«      «      }|s|||f}|�|f|z   S |S t/        ||||¬	«      S )
aö  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Binary classification labels.

        Returns:

        Examples:

        ```python
        >>> from transformers import ViltProcessor, ViltForImagesAndTextClassification
        >>> import requests
        >>> from PIL import Image

        >>> image1 = Image.open(requests.get("https://lil.nlp.cornell.edu/nlvr/exs/ex0_0.jpg", stream=True).raw)
        >>> image2 = Image.open(requests.get("https://lil.nlp.cornell.edu/nlvr/exs/ex0_1.jpg", stream=True).raw)
        >>> text = "The left image contains twice the number of dogs as the right image."

        >>> processor = ViltProcessor.from_pretrained("dandelin/vilt-b32-finetuned-nlvr2")
        >>> model = ViltForImagesAndTextClassification.from_pretrained("dandelin/vilt-b32-finetuned-nlvr2")

        >>> # prepare inputs
        >>> encoding = processor([image1, image2], text, return_tensors="pt")

        >>> # forward pass
        >>> outputs = model(input_ids=encoding.input_ids, pixel_values=encoding.pixel_values.unsqueeze(0))
        >>> logits = outputs.logits
        >>> idx = logits.argmax(-1).item()
        >>> print("Predicted answer:", model.config.id2label[idx])
        Predicted answer: True
        ```Né   r   r
   z\Make sure to match the number of images in the model with the number of images in the input.)r¢   rœ   rs   rt   rí   r�   r£   r¤   rî   r>  r?  rN   rH   r“  )rB   rî   r>  ro  ÚndimÚ	unsqueezerW   r¹  rÑ   r1  rH  rn  rr   r    r!   r&   ra   r¬  r	   rg   rR   r`   r©  r   )rC   r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  r¹  Úpooler_outputsr    r!   r“   r÷   rn  ry  r   r   r˜  r  s                            r+   r©   z*ViltForImagesAndTextClassification.forward  s’  € ð^ 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ#¨×(9Ñ(9¸QÒ(>à'×1Ñ1°!Ó4ˆLàÐ#¨×(9Ñ(9¸QÒ(>à'×1Ñ1°!Ó4ˆLà.:Ð.F�\×'Ñ'¨Ò*ÈDˆ
ØÐØ2>Ð2J˜×+Ñ+¨AÒ.ÐPTˆJØ˜Ÿ™×/Ñ/Ò/ÜØnóð ð ˆÙ2™¸ˆÙ,‘R°$ˆ
Ü�zÓ"ò 	6ˆAà—i‘iØØ-Ø-Ø<HÐ<T˜\ª!¨Q²²1²a¨-Ò8ÐZ^Ø5?Ð5K˜:¢a¨ªAªq jÒ1ÐQUØ#Ø+Ø9EÐ9Q˜\ª!¨Q²²1¨*Ò5ÐW[Ø%&¨¡UØ"3Ø%9Ø'ð  ó ˆGñ 6A˜G×1Ò1ÀgÈaÁjˆMØ×!Ñ! -Ô0Ù#Ø×$Ñ$ W×%:Ñ%:Ô;Ú Ø×!Ñ! '×"4Ñ"4Õ5ð+	6ô. Ÿ	™	 .°bÔ9ˆØ—‘ Ó/ˆàˆØÐÜ'Ó)ˆHà—Y‘Y˜vŸ}™}Ó-ˆFÙ˜FŸK™K¨¨D¯O©OÓ<¸f¿k¹kÈ"»oÓNˆDáØ˜m¨ZÐ8ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä7ØØØ'Ø!ô	
ð 	
r*   rz  )r"   r#   r$   r1   r   r{  r   r   r|  r   r&   r}  r'   r~  r   r   r©   r«   r¬   s   @r+   r·  r·  ï  so  ø„ ôñ$ +Ð+@ÓAÙÐ+SÐbqÔrð 15Ø6:Ø59Ø48Ø15Ø15Ø59Ø48Ø-1Ø,0Ø/3Ø&*ño
à˜E×,Ñ,Ñ-ðo
ð ! ×!2Ñ!2Ñ3ðo
ð ! ×!1Ñ!1Ñ2ð	o
ð
 ˜u×0Ñ0Ñ1ðo
ð ˜U×-Ñ-Ñ.ðo
ð ˜E×-Ñ-Ñ.ðo
ð   × 1Ñ 1Ñ2ðo
ð ˜u×0Ñ0Ñ1ðo
ð ˜×)Ñ)Ñ*ðo
ð $ D™>ðo
ð ' t™nðo
ð ˜d‘^ðo
ð 
Ð7¸¸u×?PÑ?PÑ9QÐQÑ	Ròo
ó só Bôo
r*   r·  zµ
    ViLT Model with a token classification head on top (a linear layer on top of the final hidden-states of the text
    tokens) e.g. for Named-Entity-Recognition (NER) tasks.
    c                    óÄ  ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 	 	 	 	 	 	 	 	 	 	 	 dde	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     de	e
j                     d	e	e
j                     d
e	e
j                     de	e
j                     de	e   de	e   de	e   deeee
j                     f   fd„«       «       Zˆ xZS )ÚViltForTokenClassificationc                 ó0  •— t         ‰| �  |«       |j                  | _        t        |d¬«      | _        t        j                  |j                  «      | _        t        j                  |j                  |j                  «      | _        | j                  «        y )NF)r`  )r0   r1   r©  rX  rH  r   r?   r@   rA   rÝ   r6   r¬  r_  rÀ   s     €r+   r1   z#ViltForTokenClassification.__init__„  sk   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ˜f¸Ô>ˆŒ	ä—z‘z &×"<Ñ"<Ó=ˆŒÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒð 	�‰Õr*   rk  r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  rû   c                 ó4  — |�|n| j                   j                  }| j                  |||||||||
||¬«      }|d   }|�|j                  d   n|j                  d   }| j	                  |«      }| j                  |dd…d|…f   «      }d}|	�Wt        «       }|	j                  |j                  «      }	 ||j                  d| j                  «      |	j                  d«      «      }|s|f|dd z   }|�|f|z   S |S t        |||j                  |j                  ¬«      S )zò
        labels (`torch.LongTensor` of shape `(batch_size, text_sequence_length)`, *optional*):
            Labels for computing the token classification loss. Indices should be in `[0, ..., config.num_labels - 1]`.

        Returns:
        Nr’  r   r   rN   rF   r“  )rB   ro  rH  rW   rA   r¬  r	   rg   rR   r`   r©  r   r    r!   )rC   r›   r¢   rœ   rs   rt   rí   r�   r£   r�  rî   r>  r?  r÷   rx  Útext_input_sizer   r   r˜  r  s                       r+   r©   z"ViltForTokenClassification.forward�  s?  € ð0 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø)Ø%Ø!ØØ'Ø%Ø/Ø!5Ø#ð ó 
ˆð " !™*ˆà09Ð0E˜)Ÿ/™/¨!Ò,È=×K^ÑK^Ð_`ÑKaˆàŸ,™, Ó7ˆØ—‘ ²Ð4D°_Ð4DÐ1DÑ!EÓFˆàˆØÐÜ'Ó)ˆHà—Y‘Y˜vŸ}™}Ó-ˆFÙ˜FŸK™K¨¨D¯O©OÓ<¸f¿k¹kÈ"»oÓNˆDáØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä$ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r*   rz  )r"   r#   r$   r1   r   r{  r   r   r|  r   r&   r}  r'   r~  r   r   r©   r«   r¬   s   @r+   rÀ  rÀ  |  s_  ø„ ô
ñ +Ð+@ÓAÙÐ+@ÈÔ_ð 15Ø6:Ø59Ø48Ø15Ø15Ø59Ø48Ø-1Ø,0Ø/3Ø&*ñ=
à˜E×,Ñ,Ñ-ð=
ð ! ×!2Ñ!2Ñ3ð=
ð ! ×!1Ñ!1Ñ2ð	=
ð
 ˜u×0Ñ0Ñ1ð=
ð ˜U×-Ñ-Ñ.ð=
ð ˜E×-Ñ-Ñ.ð=
ð   × 1Ñ 1Ñ2ð=
ð ˜u×0Ñ0Ñ1ð=
ð ˜×)Ñ)Ñ*ð=
ð $ D™>ð=
ð ' t™nð=
ð ˜d‘^ð=
ð 
Ð$ e¨E×,=Ñ,=Ñ&>Ð>Ñ	?ò=
ó `ó Bô=
r*   rÀ  )r±  r·  rÀ  r†  r§  r#  rX  rG  )Er%   Úcollections.abcrÌ   ré   Údataclassesr   Útypingr   r   r   r   r&   Útorch.utils.checkpointr   Útorch.nnr	   Úactivationsr   Úmodeling_outputsr   r   r   r   r   r   Úmodeling_utilsr   Úpytorch_utilsr   r   r   Úutilsr   r   r   r   Úconfiguration_viltr   Ú
get_loggerr"   Úloggerr|  Ú_CHECKPOINT_FOR_DOCr   ÚModuler-   r2   r8   rÔ   rú   r  r  r  r#  r.  rG  ÚVILT_START_DOCSTRINGr{  Ú4VILT_IMAGES_AND_TEXT_CLASSIFICATION_INPUTS_DOCSTRINGrX  r]  r†  rœ  rˆ  r§  r±  r·  rÀ  Ú__all__r)   r*   r+   ú<module>rÖ     s¥  ðñ ã Û Ý !ß /Ó /ã Û Ý Ý %å !÷÷ õ .÷ñ ÷
 uÓ tÝ *ð 
ˆ×	Ñ	˜HÓ	%€à€Ø-Ð ð ô@¨{ó @ó ð@ô2W!�R—Y‘Yô W!ôt6�R—Y‘Yô 6ôr˜"Ÿ)™)ô ô>9˜Ÿ	™	ô 9ôz�R—Y‘Yô ô$�B—I‘Iô ôF�r—y‘yô ô"�—‘ô ô#�—	‘	ô #ôL2
�"—)‘)ô 2
ôj*˜/ô *ð8	Ð ð5Ð ðn58Ð 4ñp ØdØóôN
Ð#ó N
ó	ðN
ôb�—‘ô ñ ðð ó	ôC
Ð)ó C
óðC
ôL "§)¡)ô ô"�"—)‘)ô ñ, ðð óôg
Ð2ó g
óðg
ñT ðð óôZ
Ð#6ó Z
óðZ
ñz ðð 9ó	ôD
Ð)<ó D
óðD
ñN ðð óôL
Ð!4ó L
óðL
ò^	�r*   