Ë
    T^(h‘  ã                   óþ  — d Z ddlZddlZddlmZmZmZ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mZmZ ddlmZ ddlmZmZmZmZ dd	lmZmZ dd
lmZmZ ddl m!Z!m"Z"m#Z#m$Z$m%Z%m&Z& ddl'm(Z(  e$jR                  e*«      Z+dZ,dZ-g d¢Z.dZ/dZ0 G d„ dejb                  «      Z2 G d„ dejb                  «      Z3	 d=dejb                  dejh                  dejh                  dejh                  deejh                     de5de5fd„Z6 G d„ d ejb                  «      Z7 G d!„ d"ejb                  «      Z8 G d#„ d$ejb                  «      Z9 G d%„ d&ejb                  «      Z: G d'„ d(ejb                  «      Z; G d)„ d*ejb                  «      Z< G d+„ d,ejb                  «      Z= G d-„ d.e«      Z>d/Z?d0Z@ e"d1e?«       G d2„ d3e>«      «       ZA G d4„ d5ejb                  «      ZB e"d6e?«       G d7„ d8e>«      «       ZC e"d9e?«       G d:„ d;e>«      «       ZDg d<¢ZEy)>zPyTorch ViT model.é    N)ÚCallableÚDictÚListÚOptionalÚSetÚTupleÚUnion)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )ÚACT2FN)ÚBaseModelOutputÚBaseModelOutputWithPoolingÚImageClassifierOutputÚMaskedImageModelingOutput)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)Ú find_pruneable_heads_and_indicesÚprune_linear_layer)Úadd_code_sample_docstringsÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstringsÚ	torch_inté   )Ú	ViTConfigr   z!google/vit-base-patch16-224-in21k)r   éÅ   i   zgoogle/vit-base-patch16-224zEgyptian catc            	       óÒ   ‡ — 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ej                     dedej                  fd„Zˆ xZS )ÚViTEmbeddingszb
    Construct the CLS token, position and patch embeddings. Optionally, also the mask token.
    ÚconfigÚuse_mask_tokenÚreturnNc                 óJ  •— t         ‰| �  «        t        j                  t	        j
                  dd|j                  «      «      | _        |r4t        j                  t	        j                  dd|j                  «      «      nd | _	        t        |«      | _        | j                  j                  }t        j                  t	        j
                  d|dz   |j                  «      «      | _        t        j                  |j                  «      | _        |j"                  | _        || _        y )Nr   )ÚsuperÚ__init__r
   Ú	ParameterÚtorchÚrandnÚhidden_sizeÚ	cls_tokenÚzerosÚ
mask_tokenÚViTPatchEmbeddingsÚpatch_embeddingsÚnum_patchesÚposition_embeddingsÚDropoutÚhidden_dropout_probÚdropoutÚ
patch_sizer#   )Úselfr#   r$   r2   Ú	__class__s       €úb/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/vit/modeling_vit.pyr(   zViTEmbeddings.__init__A   sÊ   ø€ Ü‰ÑÔäŸ™¤e§k¡k°!°Q¸×8JÑ8JÓ&KÓLˆŒÙQ_œ"Ÿ,™,¤u§{¡{°1°a¸×9KÑ9KÓ'LÔMÐeiˆŒÜ 2°6Ó :ˆÔØ×+Ñ+×7Ñ7ˆÜ#%§<¡<´·±¸A¸{ÈQ¹ÐPV×PbÑPbÓ0cÓ#dˆÔ Ü—z‘z &×"<Ñ"<Ó=ˆŒØ ×+Ñ+ˆŒØˆ�ó    Ú
embeddingsÚheightÚwidthc                 ó¦  — |j                   d   dz
  }| j                  j                   d   dz
  }t        j                  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  «      }|j                  d|||«      }|j                  dddd«      }t        j                  j                  ||	|
fdd	¬
«      }|j                  dddd«      j                  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éÿÿÿÿç      à?r   r   é   ÚbicubicF)ÚsizeÚmodeÚalign_corners©Údim)Úshaper3   r*   ÚjitÚ
is_tracingr7   r   ÚreshapeÚpermuter
   Ú
functionalÚinterpolateÚviewÚcat)r8   r<   r=   r>   r2   Únum_positionsÚclass_pos_embedÚpatch_pos_embedrH   Ú
new_heightÚ	new_widthÚsqrt_num_positionss               r:   Úinterpolate_pos_encodingz&ViTEmbeddings.interpolate_pos_encodingM   s`  € ð !×&Ñ& qÑ)¨AÑ-ˆØ×0Ñ0×6Ñ6°qÑ9¸AÑ=ˆô �y‰y×#Ñ#Ô%¨+¸Ò*FÈ6ÐUZÊ?Ø×+Ñ+Ð+à×2Ñ2²1°b°q°b°5Ñ9ˆØ×2Ñ2²1°a±b°5Ñ9ˆà×Ñ˜rÑ"ˆà˜tŸ™Ñ.ˆ
Ø˜TŸ_™_Ñ,ˆ	ä& }°cÑ'9Ó:ÐØ)×1Ñ1°!Ð5GÐI[Ð]`ÓaˆØ)×1Ñ1°!°Q¸¸1Ó=ˆäŸ-™-×3Ñ3ØØ˜iÐ(ØØð	 4ó 
ˆð *×1Ñ1°!°Q¸¸1Ó=×BÑBÀ1ÀbÈ#ÓNˆä�y‰y˜/¨?Ð;ÀÔCÐCr;   Úpixel_valuesÚbool_masked_posrX   c                 óä  — |j                   \  }}}}| j                  ||¬«      }|�Z|j                   d   }	| j                  j                  ||	d«      }
|j	                  d«      j                  |
«      }|d|z
  z  |
|z  z   }| j                  j                  |dd«      }t        j                  ||fd¬«      }|r|| j                  |||«      z   }n|| j                  z   }| j                  |«      }|S )N)rX   r   r@   ç      ð?rG   )rI   r1   r/   ÚexpandÚ	unsqueezeÚtype_asr-   r*   rQ   rX   r3   r6   )r8   rY   rZ   rX   Ú
batch_sizeÚnum_channelsr=   r>   r<   Ú
seq_lengthÚmask_tokensÚmaskÚ
cls_tokenss                r:   ÚforwardzViTEmbeddings.forwardu   s  € ð 3?×2DÑ2DÑ/ˆ
�L &¨%Ø×*Ñ*¨<ÐRjÐ*Ókˆ
àÐ&Ø#×)Ñ)¨!Ñ,ˆJØŸ/™/×0Ñ0°¸ZÈÓLˆKà"×,Ñ,¨RÓ0×8Ñ8¸ÓEˆDØ# s¨T¡zÑ2°[À4Ñ5GÑGˆJð —^‘^×*Ñ*¨:°r¸2Ó>ˆ
Ü—Y‘Y 
¨JÐ7¸QÔ?ˆ
ñ $Ø# d×&CÑ&CÀJÐPVÐX]Ó&^Ñ^‰Jà# d×&>Ñ&>Ñ>ˆJà—\‘\ *Ó-ˆ
àÐr;   ©F©NF)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úboolr(   r*   ÚTensorÚintrX   r   Ú
BoolTensorrf   Ú__classcell__©r9   s   @r:   r"   r"   <   s›   ø„ ññ
˜yð 
¸$ð 
È4õ 
ð&D°5·<±<ð &DÈð &DÐUXð &DÐ]b×]iÑ]ió &DðV 7;Ø).ñ	à—l‘lðð " %×"2Ñ"2Ñ3ðð #'ð	ð
 
�‰÷r;   r"   c                   ó`   ‡ — e Zd ZdZˆ fd„Zddej                  dedej                  fd„Zˆ xZ	S )r0   zì
    This class turns `pixel_values` of shape `(batch_size, num_channels, height, width)` into the initial
    `hidden_states` (patch embeddings) of shape `(batch_size, seq_length, hidden_size)` to be consumed by a
    Transformer.
    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)r'   r(   Ú
image_sizer7   ra   r,   Ú
isinstanceÚcollectionsÚabcÚIterabler2   r
   ÚConv2dÚ
projection)r8   r#   rw   r7   ra   r,   r2   r9   s          €r:   r(   zViTPatchEmbeddings.__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;   rY   rX   r%   c                 óŽ  — |j                   \  }}}}|| j                  k7  rt        d| j                  › d|› d�«      ‚|sV|| j                  d   k7  s|| j                  d   k7  r2t        d|› d|› d| j                  d   › d| j                  d   › d	�	«      ‚| j	                  |«      j                  d
«      j                  dd
«      }|S )NzoMake sure that the channel dimension of the pixel values match with the one set in the configuration. Expected z	 but got ú.r   r   zInput image size (Ú*z) doesn't match model (z).rB   )rI   ra   Ú
ValueErrorrw   r}   ÚflattenÚ	transpose)r8   rY   rX   r`   ra   r=   r>   r<   s           r:   rf   zViTPatchEmbeddings.forwardª   sê   € Ø2>×2DÑ2DÑ/ˆ
�L &¨%Ø˜4×,Ñ,Ò,ÜðØ!×.Ñ.Ð/¨y¸¸ÀaðIóð ñ (Ø˜Ÿ™¨Ñ+Ò+¨u¸¿¹ÈÑ8JÒ/JÜ Ø(¨¨°°%°ð 9ØŸ™¨Ñ+Ð,¨A¨d¯o©o¸aÑ.@Ð-AÀðEóð ð —_‘_ \Ó2×:Ñ:¸1Ó=×GÑGÈÈ1ÓMˆ
ØÐr;   rg   )
ri   rj   rk   rl   r(   r*   rn   rm   rf   rq   rr   s   @r:   r0   r0   ”   s3   ø„ ñôjñ E§L¡Lð ÈDð Ð]b×]iÑ]i÷ r;   r0   ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingr6   c                 óÀ  — t        j                  ||j                  dd«      «      |z  }t        j                  j                  |dt         j                  ¬«      j                  |j                  «      }t        j                  j                  ||| j                  ¬«      }|�||z  }t        j                  ||«      }	|	j                  dd«      j                  «       }	|	|fS )Nr@   éþÿÿÿ)rH   Údtype)ÚpÚtrainingr   rB   )r*   Úmatmulrƒ   r
   rN   ÚsoftmaxÚfloat32ÚtorŒ   r6   rŽ   Ú
contiguous)
r„   r…   r†   r‡   rˆ   r‰   r6   ÚkwargsÚattn_weightsÚattn_outputs
             r:   Úeager_attention_forwardr—   »   sÀ   € ô —<‘<  s§}¡}°R¸Ó'<Ó=ÀÑG€Lô —=‘=×(Ñ(¨¸2ÄUÇ]Á]Ð(ÓS×VÑVÐW\×WbÑWbÓc€Lô —=‘=×(Ñ(¨¸È6Ï?É?Ð(Ó[€Lð Ð!Ø# nÑ4ˆä—,‘,˜|¨UÓ3€KØ×'Ñ'¨¨1Ó-×8Ñ8Ó:€Kà˜Ð$Ð$r;   c            
       óè   ‡ — e Zd Zdeddfˆ fd„Zdej                  dej                  fd„Z	 d
deej                     de	de
eej                  ej                  f   eej                     f   fd	„Zˆ xZS )ÚViTSelfAttentionr#   r%   Nc                 ó2  •— t         ‰| �  «        |j                  |j                  z  dk7  r2t	        |d«      s&t        d|j                  › d|j                  › d�«      ‚|| _        |j                  | _        t        |j                  |j                  z  «      | _        | j                  | j                  z  | _	        |j                  | _        | j                  dz  | _        d| _        t        j                  |j                  | j                  |j                   ¬«      | _        t        j                  |j                  | j                  |j                   ¬«      | _        t        j                  |j                  | j                  |j                   ¬«      | _        y )	Nr   Úembedding_sizezThe hidden size z4 is not a multiple of the number of attention heads r   g      à¿F)Úbias)r'   r(   r,   Únum_attention_headsÚhasattrr�   r#   ro   Úattention_head_sizeÚall_head_sizeÚattention_probs_dropout_probÚdropout_probr‰   Ú	is_causalr
   ÚLinearÚqkv_biasr…   r†   r‡   ©r8   r#   r9   s     €r:   r(   zViTSelfAttention.__init__Ú   sF  ø€ Ü‰ÑÔØ×Ñ × :Ñ :Ñ:¸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ˆÔØ"×?Ñ?ˆÔØ×/Ñ/°Ñ5ˆŒØˆŒä—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Ô\ˆ�
r;   Úxc                 ó¤   — |j                  «       d d | j                  | j                  fz   }|j                  |«      }|j	                  dddd«      S )Nr@   r   rB   r   r   )rD   r�   rŸ   rP   rM   )r8   r§   Únew_x_shapes      r:   Útranspose_for_scoresz%ViTSelfAttention.transpose_for_scoresî   sL   € Ø—f‘f“h˜s �m t×'?Ñ'?À×AYÑAYÐ&ZÑZˆØ�F‰F�;ÓˆØ�y‰y˜˜A˜q !Ó$Ð$r;   Ú	head_maskÚoutput_attentionsc           
      ó˜  — | j                  | j                  |«      «      }| j                  | j                  |«      «      }| j                  | j                  |«      «      }t        }| j
                  j                  dk7  rN| j
                  j                  dk(  r|rt        j                  d«       nt        | j
                  j                     } || ||||| j                  | j                  | j                  sdn| j                  ¬«      \  }}	|j                  «       d d | j                  fz   }
|j!                  |
«      }|r||	f}|S |f}|S )NÚeagerÚsdpazã`torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to eager attention. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.ç        )r£   r‰   r6   r‹   )rª   r†   r‡   r…   r—   r#   Ú_attn_implementationÚloggerÚwarning_oncer   r£   r‰   rŽ   r¢   rD   r    rL   )r8   Úhidden_statesr«   r¬   Ú	key_layerÚvalue_layerÚquery_layerÚattention_interfaceÚcontext_layerÚattention_probsÚnew_context_layer_shapeÚoutputss               r:   rf   zViTSelfAttention.forwardó   s=  € ð ×-Ñ-¨d¯h©h°}Ó.EÓFˆ	Ø×/Ñ/°·
±
¸=Ó0IÓJˆØ×/Ñ/°·
±
¸=Ó0IÓJˆä(?ÐØ�;‰;×+Ñ+¨wÒ6Ø�{‰{×/Ñ/°6Ò9Ñ>OÜ×#Ñ#ðLõô
 '>¸d¿k¹k×>^Ñ>^Ñ&_Ð#á)<ØØØØØØ—n‘nØ—L‘LØ#Ÿ}š}‘C°$×2CÑ2Cô	*
Ñ&ˆ�ð #0×"4Ñ"4Ó"6°s¸Ð";¸t×?QÑ?QÐ>SÑ"SÐØ%×-Ñ-Ð.EÓFˆá6G�= /Ð2ˆàˆð O\ÐM]ˆàˆr;   rh   )ri   rj   rk   r   r(   r*   rn   rª   r   rm   r	   r   rf   rq   rr   s   @r:   r™   r™   Ù   s†   ø„ ð]˜yð ]¨Tõ ]ð(% e§l¡lð %°u·|±|ó %ð bgñ!Ø(0°·±Ñ(>ð!ØZ^ð!à	ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷!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 )	ÚViTSelfOutputz 
    The residual connection is defined in ViTLayer instead of here (as is the case with other models), due to the
    layernorm applied before each block.
    r#   r%   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  |j                  «      | _        y ©N)	r'   r(   r
   r¤   r,   Údenser4   r5   r6   r¦   s     €r:   r(   zViTSelfOutput.__init__  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r;   r´   Úinput_tensorc                 óJ   — | j                  |«      }| j                  |«      }|S rÀ   ©rÁ   r6   ©r8   r´   rÂ   s      r:   rf   zViTSelfOutput.forward"  s$   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆàÐr;   )
ri   rj   rk   rl   r   r(   r*   rn   rf   rq   rr   s   @r:   r¾   r¾     sD   ø„ ñð
>˜yð >¨Tõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r;   r¾   c                   óà   ‡ — e Zd Zdeddfˆ fd„Zdee   ddfd„Z	 	 ddej                  de
ej                     d	edeeej                  ej                  f   eej                     f   fd
„Zˆ xZS )ÚViTAttentionr#   r%   Nc                 ó€   •— t         ‰| �  «        t        |«      | _        t	        |«      | _        t        «       | _        y rÀ   )r'   r(   r™   Ú	attentionr¾   ÚoutputÚsetÚpruned_headsr¦   s     €r:   r(   zViTAttention.__init__*  s0   ø€ Ü‰ÑÔÜ)¨&Ó1ˆŒÜ# FÓ+ˆŒÜ›EˆÕr;   Úheadsc                 ó>  — 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   rG   )Úlenr   rÉ   r�   rŸ   rÌ   r   r…   r†   r‡   rÊ   rÁ   r    Úunion)r8   rÍ   Úindexs      r:   Úprune_headszViTAttention.prune_heads0  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;   r´   r«   r¬   c                 óh   — | j                  |||«      }| j                  |d   |«      }|f|dd  z   }|S )Nr   r   )rÉ   rÊ   )r8   r´   r«   r¬   Úself_outputsÚattention_outputr¼   s          r:   rf   zViTAttention.forwardB  sE   € ð —~‘~ m°YÐ@QÓRˆàŸ;™; |°A¡¸ÓFÐà#Ð%¨°Q°RÐ(8Ñ8ˆØˆr;   rh   )ri   rj   rk   r   r(   r   ro   rÒ   r*   rn   r   rm   r	   r   rf   rq   rr   s   @r:   rÇ   rÇ   )  s’   ø„ ð"˜yð "¨Tõ "ð;  S¡ð ;¨dó ;ð* -1Ø"'ñ	à—|‘|ðð ˜EŸL™LÑ)ðð  ð	ð
 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷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 )ÚViTIntermediater#   r%   Nc                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y rÀ   )r'   r(   r
   r¤   r,   Úintermediate_sizerÁ   rx   Ú
hidden_actÚstrr   Úintermediate_act_fnr¦   s     €r:   r(   zViTIntermediate.__init__Q  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Ü   )r8   r´   s     r:   rf   zViTIntermediate.forwardY  s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆàÐr;   ©	ri   rj   rk   r   r(   r*   rn   rf   rq   rr   s   @r:   r×   r×   P  s1   ø„ ð9˜yð 9¨Tõ 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 )Ú	ViTOutputr#   r%   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j                  «      | _	        y rÀ   )
r'   r(   r
   r¤   rÙ   r,   rÁ   r4   r5   r6   r¦   s     €r:   r(   zViTOutput.__init__a  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:   rf   zViTOutput.forwardf  s.   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆà%¨Ñ4ˆàÐr;   rÞ   rr   s   @r:   rà   rà   `  s?   ø„ ð>˜yð >¨Tõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r;   rà   c                   óÎ   ‡ — e Zd ZdZdeddfˆ fd„Z	 	 d
dej                  deej                     de	de
eej                  ej                  f   eej                     f   fd	„Zˆ xZS )ÚViTLayerz?This corresponds to the Block class in the timm implementation.r#   r%   Nc                 ór  •— t         ‰| �  «        |j                  | _        d| _        t	        |«      | _        t        |«      | _        t        |«      | _	        t        j                  |j                  |j                  ¬«      | _        t        j                  |j                  |j                  ¬«      | _        y )Nr   ©Úeps)r'   r(   Úchunk_size_feed_forwardÚseq_len_dimrÇ   rÉ   r×   Úintermediaterà   rÊ   r
   Ú	LayerNormr,   Úlayer_norm_epsÚlayernorm_beforeÚlayernorm_afterr¦   s     €r:   r(   zViTLayer.__init__r  s‡   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ% fÓ-ˆŒÜ+¨FÓ3ˆÔÜ Ó'ˆŒÜ "§¡¨V×-?Ñ-?ÀV×EZÑEZÔ [ˆÔÜ!Ÿ|™|¨F×,>Ñ,>ÀF×DYÑDYÔZˆÕr;   r´   r«   r¬   c                 óÞ   — | j                  | j                  |«      ||¬«      }|d   }|dd  }||z   }| j                  |«      }| j                  |«      }| j	                  ||«      }|f|z   }|S )N)r¬   r   r   )rÉ   rí   rî   rê   rÊ   )r8   r´   r«   r¬   Úself_attention_outputsrÕ   r¼   Úlayer_outputs           r:   rf   zViTLayer.forward|  s–   € ð "&§¡Ø×!Ñ! -Ó0ØØ/ð "0ó "
Ðð
 2°!Ñ4ÐØ(¨¨Ð,ˆð )¨=Ñ8ˆð ×+Ñ+¨MÓ:ˆØ×(Ñ(¨Ó6ˆð —{‘{ <°Ó?ˆà�/ GÑ+ˆàˆr;   rh   )ri   rj   rk   rl   r   r(   r*   rn   r   rm   r	   r   rf   rq   rr   s   @r:   rä   rä   o  s�   ø„ ÙIð[˜yð [¨Tõ [ð -1Ø"'ñ	à—|‘|ðð ˜EŸL™LÑ)ðð  ð	ð
 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷r;   rä   c                   óŠ   ‡ — e Zd Zdeddfˆ fd„Z	 	 	 	 ddej                  deej                     deded	ede	e
ef   fd
„Zˆ xZS )Ú
ViTEncoderr#   r%   Nc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w rh   )
r'   r(   r#   r
   Ú
ModuleListÚrangeÚnum_hidden_layersrä   ÚlayerÚgradient_checkpointing)r8   r#   Ú_r9   s      €r:   r(   zViTEncoder.__init__š  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]¼eÀF×D\ÑD\Ó>]Ö#^¸¤H¨VÕ$4Ò#^Ó_ˆŒ
Ø&+ˆÕ#ùò $_s   ½A#r´   r«   r¬   Úoutput_hidden_statesÚreturn_dictc                 ót  — |rdnd }|rdnd }t        | j                  «      D ]h  \  }}	|r||fz   }|�||   nd }
| j                  r+| j                  r| j	                  |	j
                  ||
|«      }n
 |	||
|«      }|d   }|sŒ`||d   fz   }Œj |r||fz   }|st        d„ |||fD «       «      S t        |||¬«      S )N© r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrÀ   rþ   )Ú.0Úvs     r:   ú	<genexpr>z%ViTEncoder.forward.<locals>.<genexpr>Ä  s   è ø€ Òm˜qÐ_`Ñ_lœÑmùs   ‚Š)Úlast_hidden_stater´   Ú
attentions)Ú	enumeraterø   rù   rŽ   Ú_gradient_checkpointing_funcÚ__call__Útupler   )r8   r´   r«   r¬   rû   rü   Úall_hidden_statesÚall_self_attentionsÚiÚlayer_moduleÚlayer_head_maskÚlayer_outputss               r:   rf   zViTEncoder.forward   sÿ   € ñ #7™B¸DÐÙ$5™b¸4Ðä(¨¯©Ó4ò 	P‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à.7Ð.C˜i¨šlÈˆOà×*Ò*¨t¯}ª}Ø $× AÑ AØ ×)Ñ)Ø!Ø#Ø%ó	!‘ñ !-¨]¸OÐM^Ó _�à)¨!Ñ,ˆMâ Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð'	Pñ*  Ø 1°]Ð4DÑ DÐáÜÑm ]Ð4EÐGZÐ$[ÔmÓmÐmÜØ+Ø+Ø*ô
ð 	
r;   )NFFT)ri   rj   rk   r   r(   r*   rn   r   rm   r	   r  r   rf   rq   rr   s   @r:   ró   ró   ™  sz   ø„ ð,˜yð ,¨Tõ ,ð -1Ø"'Ø%*Ø ñ)
à—|‘|ð)
ð ˜EŸL™LÑ)ð)
ð  ð	)
ð
 #ð)
ð ð)
ð 
ˆu�oÐ%Ñ	&÷)
r;   ró   c                   óˆ   — e Zd ZdZeZdZdZdZddgZ	dZ
dZdeej                  ej                  ej                   f   dd	fd
„Zy	)ÚViTPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚvitrY   Tr"   rä   r„   r%   Nc                 ó   — t        |t        j                  t        j                  f«      rËt        j                  j                  |j                  j                  j                  t        j                  «      d| j                  j                  ¬«      j                  |j                  j                  «      |j                  _        |j                  �%|j                  j                  j                  «        yyt        |t        j                   «      rJ|j                  j                  j                  «        |j                  j                  j#                  d«       yt        |t$        «      �rdt        j                  j                  |j&                  j                  j                  t        j                  «      d| j                  j                  ¬«      j                  |j&                  j                  «      |j&                  _        t        j                  j                  |j(                  j                  j                  t        j                  «      d| j                  j                  ¬«      j                  |j(                  j                  «      |j(                  _        |j*                  �%|j*                  j                  j                  «        yyy)zInitialize the weightsr°   )ÚmeanÚstdNr\   )rx   r
   r¤   r|   ÚinitÚtrunc_normal_ÚweightÚdatar’   r*   r‘   r#   Úinitializer_rangerŒ   rœ   Úzero_rë   Úfill_r"   r3   r-   r/   )r8   r„   s     r:   Ú_init_weightsz ViTPreTrainedModel._init_weightsÚ  sò  € ä�fœrŸy™y¬"¯)©)Ð4Ô5ô "$§¡×!6Ñ!6Ø—‘×"Ñ"×%Ñ%¤e§m¡mÓ4¸3ÀDÇKÁK×DaÑDað "7ó "ç‰b�—‘×$Ñ$Ó%ð �M‰MÔð �{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡Ô-Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)Ü˜¤Õ.Ü.0¯g©g×.CÑ.CØ×*Ñ*×/Ñ/×2Ñ2´5·=±=ÓAØØ—K‘K×1Ñ1ð /Dó /÷ ‰b�×+Ñ+×1Ñ1Ó2ð	 ×&Ñ&Ô+ô %'§G¡G×$9Ñ$9Ø× Ñ ×%Ñ%×(Ñ(¬¯©Ó7ØØ—K‘K×1Ñ1ð %:ó %÷ ‰b�×!Ñ!×'Ñ'Ó(ð	 ×ÑÔ!ð × Ñ Ð,Ø×!Ñ!×&Ñ&×,Ñ,Õ.ð -ð /r;   )ri   rj   rk   rl   r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_supports_sdpaÚ_supports_flash_attn_2r	   r
   r¤   r|   rë   r  rþ   r;   r:   r  r  Ì  sa   „ ñð
 €LØÐØ$€OØ&*Ð#Ø(¨*Ð5ÐØ€NØ!Ðð/ E¨"¯)©)°R·Y±YÀÇÁÐ*LÑ$Mð /ÐRVô /r;   r  aE  
    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 ([`ViTConfig`]): 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:
        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See [`ViTImageProcessor.__call__`]
            for details.

        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**.

        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.
        interpolate_pos_encoding (`bool`, *optional*):
            Whether to interpolate the pre-trained position encodings.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
z]The bare ViT Model transformer outputting raw hidden-states without any specific head on top.c                   ó8  ‡ — e Zd Zddededefˆ fd„Zdefd„Zdee	e
e	   f   ddfd	„Z ee«       eeeed
e¬«      	 	 	 	 	 	 	 d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f   fd„«       «       Zˆ xZS )ÚViTModelr#   Úadd_pooling_layerr$   c                 ó  •— t         ‰| �  |«       || _        t        ||¬«      | _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _        |rt        |«      nd | _        | j                  «        y )N)r$   ræ   )r'   r(   r#   r"   r<   ró   Úencoderr
   rë   r,   rì   Ú	layernormÚ	ViTPoolerÚpoolerÚ	post_init)r8   r#   r&  r$   r9   s       €r:   r(   zViTModel.__init__!  sk   ø€ Ü‰Ñ˜Ô ØˆŒä'¨¸~ÔNˆŒÜ! &Ó)ˆŒäŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÙ+<”i Ô'À$ˆŒð 	�‰Õr;   r%   c                 ó.   — | j                   j                  S rÀ   )r<   r1   )r8   s    r:   Úget_input_embeddingszViTModel.get_input_embeddings.  s   € Ø�‰×/Ñ/Ð/r;   Úheads_to_pruneNc                 ó˜   — |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(  rø   rÉ   rÒ   )r8   r/  rø   rÍ   s       r:   Ú_prune_headszViTModel._prune_heads1  sE   € ð
 +×0Ñ0Ó2ò 	C‰LˆE�5Ø�L‰L×Ñ˜uÑ%×/Ñ/×;Ñ;¸EÕBñ	Cr;   Úvision)Ú
checkpointÚoutput_typer  ÚmodalityÚexpected_outputrY   rZ   r«   r¬   rû   rX   rü   c                 óÖ  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }|€t	        d«      ‚| j                  || j                   j                  «      }| j                  j                  j                  j                  j                  }|j                  |k7  r|j                  |«      }| j                  |||¬«      }	| j                  |	||||¬«      }
|
d   }| j                  |«      }| j                  �| j                  |«      nd}|s|�||fn|f}||
dd z   S t!        |||
j"                  |
j$                  ¬«      S )zË
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, num_patches)`, *optional*):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).
        Nz You have to specify pixel_values)rZ   rX   )r«   r¬   rû   rü   r   r   )r  Úpooler_outputr´   r  )r#   r¬   rû   Úuse_return_dictr�   Úget_head_maskr÷   r<   r1   r}   r  rŒ   r’   r(  r)  r+  r   r´   r  )r8   rY   rZ   r«   r¬   rû   rX   rü   Úexpected_dtypeÚembedding_outputÚencoder_outputsÚsequence_outputÚpooled_outputÚhead_outputss                 r:   rf   zViTModel.forward9  s„  € ð, 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐÜÐ?Ó@Ð@ð ×&Ñ& y°$·+±+×2OÑ2OÓPˆ	ð Ÿ™×9Ñ9×DÑD×KÑK×QÑQˆØ×Ñ Ò/Ø'Ÿ?™?¨>Ó:ˆLàŸ?™?Ø¨/ÐTlð +ó 
Ðð Ÿ,™,ØØØ/Ø!5Ø#ð 'ó 
ˆð *¨!Ñ,ˆØŸ.™.¨Ó9ˆØ8<¿¹Ð8O˜Ÿ™ OÔ4ÐUYˆáØ?LÐ?X˜O¨]Ñ;Ð_nÐ^pˆLØ /°!°"Ð"5Ñ5Ð5ä)Ø-Ø'Ø)×7Ñ7Ø&×1Ñ1ô	
ð 	
r;   )TF©NNNNNNN)ri   rj   rk   r   rm   r(   r0   r.  r   ro   r   r2  r   ÚVIT_INPUTS_DOCSTRINGr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCÚ_EXPECTED_OUTPUT_SHAPEr   r*   rn   rp   r	   r   rf   rq   rr   s   @r:   r%  r%    s"  ø„ ñ
˜yð ¸Tð ÐZ^õ ð0Ð&8ó 0ðC¨4°°T¸#±Y°Ñ+?ð CÀDó Cñ +Ð+?Ó@ÙØ&Ø.Ø$ØØ.ôð 04Ø6:Ø,0Ø,0Ø/3Ø37Ø&*ñ;
à˜uŸ|™|Ñ,ð;
ð " %×"2Ñ"2Ñ3ð;
ð ˜EŸL™LÑ)ð	;
ð
 $ D™>ð;
ð ' t™nð;
ð #+¨4¡.ð;
ð ˜d‘^ð;
ð 
ˆuÐ0Ð0Ñ	1ò;
óó Aô;
r;   r%  c                   ó*   ‡ — e Zd Zdefˆ fd„Zd„ Zˆ xZS )r*  r#   c                 ó°   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                     | _	        y rÀ   )
r'   r(   r
   r¤   r,   Úpooler_output_sizerÁ   r   Ú
pooler_actÚ
activationr¦   s     €r:   r(   zViTPooler.__init__€  s>   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3LÑ3LÓMˆŒ
Ü  ×!2Ñ!2Ñ3ˆ�r;   c                 ó\   — |d d …df   }| j                  |«      }| j                  |«      }|S )Nr   )rÁ   rK  )r8   r´   Úfirst_token_tensorr@  s       r:   rf   zViTPooler.forward…  s6   € ð +ª1¨a¨4Ñ0ÐØŸ
™
Ð#5Ó6ˆØŸ™¨Ó6ˆØÐr;   )ri   rj   rk   r   r(   rf   rq   rr   s   @r:   r*  r*    s   ø„ ð4˜yõ 4ö
r;   r*  aV  ViT Model with a decoder on top for masked image modeling, as proposed in [SimMIM](https://arxiv.org/abs/2111.09886).

    <Tip>

    Note that we provide a script to pre-train this model on custom data in our [examples
    directory](https://github.com/huggingface/transformers/tree/main/examples/pytorch/image-pretraining).

    </Tip>
    c                   ó   ‡ — e Zd Zdeddfˆ fd„Z ee«       eee	¬«      	 	 	 	 	 	 	 d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f   fd„«       «       Zˆ xZS )ÚViTForMaskedImageModelingr#   r%   Nc                 óN  •— t         ‰| �  |«       t        |dd¬«      | _        t	        j
                  t	        j                  |j                  |j                  dz  |j                  z  d¬«      t	        j                  |j                  «      «      | _        | j                  «        y )NFT)r&  r$   rB   r   )Úin_channelsÚout_channelsru   )r'   r(   r%  r  r
   Ú
Sequentialr|   r,   Úencoder_stridera   ÚPixelShuffleÚdecoderr,  r¦   s     €r:   r(   z"ViTForMaskedImageModeling.__init__›  s�   ø€ Ü‰Ñ˜Ô ä˜F°eÈDÔQˆŒä—}‘}Ü�I‰IØ"×.Ñ.Ø#×2Ñ2°AÑ5¸×8KÑ8KÑKØôô
 �O‰O˜F×1Ñ1Ó2ó
ˆŒð 	�‰Õr;   )r5  r  rY   rZ   r«   r¬   rû   rX   rü   c           	      ó   — |�|n| j                   j                  }|�g| j                   j                  | j                   j                  k7  r:t	        d| j                   j                  › d| j                   j                  › d�«      ‚| j                  |||||||¬«      }|d   }	|	dd…dd…f   }	|	j                  \  }
}}t        j                  |dz  «      x}}|	j                  dd	d«      j                  |
|||«      }	| j                  |	«      }d}|��| j                   j                  | j                   j                  z  }|j                  d
||«      }|j                  | j                   j                  d«      j                  | j                   j                  d	«      j                  d«      j                  «       }t         j"                  j%                  ||d¬«      }||z  j'                  «       |j'                  «       dz   z  | j                   j(                  z  }|s|f|dd z   }|�|f|z   S |S t+        |||j,                  |j.                  ¬«      S )a=  
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, num_patches)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).

        Returns:

        Examples:
        ```python
        >>> from transformers import AutoImageProcessor, ViTForMaskedImageModeling
        >>> import torch
        >>> from PIL import Image
        >>> import requests

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

        >>> image_processor = AutoImageProcessor.from_pretrained("google/vit-base-patch16-224-in21k")
        >>> model = ViTForMaskedImageModeling.from_pretrained("google/vit-base-patch16-224-in21k")

        >>> num_patches = (model.config.image_size // model.config.patch_size) ** 2
        >>> pixel_values = image_processor(images=image, return_tensors="pt").pixel_values
        >>> # create random boolean mask of shape (batch_size, num_patches)
        >>> bool_masked_pos = torch.randint(low=0, high=2, size=(1, num_patches)).bool()

        >>> outputs = model(pixel_values, bool_masked_pos=bool_masked_pos)
        >>> loss, reconstructed_pixel_values = outputs.loss, outputs.reconstruction
        >>> list(reconstructed_pixel_values.shape)
        [1, 3, 224, 224]
        ```Nz³When `bool_masked_pos` is provided, `patch_size` must be equal to `encoder_stride` to ensure that the reconstructed image has the same dimensions as the input. Got `patch_size` = z and `encoder_stride` = r   )rZ   r«   r¬   rû   rX   rü   r   r   rA   rB   r@   Únone)Ú	reductiongñhãˆµøä>)ÚlossÚreconstructionr´   r  )r#   r:  r7   rT  r�   r  rI   ÚmathÚfloorrM   rL   rV  rw   Úrepeat_interleaver^   r“   r
   rN   Úl1_lossÚsumra   r   r´   r  )r8   rY   rZ   r«   r¬   rû   rX   rü   r¼   r?  r`   Úsequence_lengthra   r=   r>   Úreconstructed_pixel_valuesÚmasked_im_lossrD   rd   Úreconstruction_lossrÊ   s                        r:   rf   z!ViTForMaskedImageModeling.forward¬  sR  € ðR &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ&¨D¯K©K×,BÑ,BÀdÇkÁk×F`ÑF`Ò,`Üð&à&*§k¡k×&<Ñ&<Ð%=Ð=UÐVZ×VaÑVa×VpÑVpÐUqÐqrðtóð ð —(‘(ØØ+ØØ/Ø!5Ø%=Ø#ð ó 
ˆð " !™*ˆð *ª!¨Q©R¨%Ñ0ˆØ4C×4IÑ4IÑ1ˆ
�O \ÜŸ™ O°SÑ$8Ó9Ð9ˆ�Ø)×1Ñ1°!°Q¸Ó:×BÑBÀ:È|Ð]cÐejÓkˆð &*§\¡\°/Ó%BÐ"àˆØÑ&Ø—;‘;×)Ñ)¨T¯[©[×-CÑ-CÑCˆDØ-×5Ñ5°b¸$ÀÓEˆOà×1Ñ1°$·+±+×2HÑ2HÈ!ÓLß"Ñ" 4§;¡;×#9Ñ#9¸1Ó=ß‘˜1“ß‘“ð	 ô #%§-¡-×"7Ñ"7¸ÐF`ÐlrÐ"7Ó"sÐØ1°DÑ8×=Ñ=Ó?À4Ç8Á8Ã:ÐPTÑCTÑUÐX\×XcÑXc×XpÑXpÑpˆNáØ0Ð2°W¸Q¸R°[Ñ@ˆFØ3AÐ3M�^Ð%¨Ñ.ÐYÐSYÐYä(ØØ5Ø!×/Ñ/Ø×)Ñ)ô	
ð 	
r;   rB  )ri   rj   rk   r   r(   r   rC  r   r   rE  r   r*   rn   rp   rm   r	   r  rf   rq   rr   s   @r:   rO  rO  Ž  sê   ø„ ð˜yð ¨Tõ ñ" +Ð+?Ó@ÙÐ+DÐSbÔcð 04Ø6:Ø,0Ø,0Ø/3Ø37Ø&*ñ[
à˜uŸ|™|Ñ,ð[
ð " %×"2Ñ"2Ñ3ð[
ð ˜EŸL™LÑ)ð	[
ð
 $ D™>ð[
ð ' t™nð[
ð #+¨4¡.ð[
ð ˜d‘^ð[
ð 
ˆuÐ/Ð/Ñ	0ò[
ó dó Aô[
r;   rO  aà  
    ViT Model transformer with an image classification head on top (a linear layer on top of the final hidden state of
    the [CLS] token) e.g. for ImageNet.

    <Tip>

        Note that it's possible to fine-tune ViT on higher resolution images than the ones it has been trained on, by
        setting `interpolate_pos_encoding` to `True` in the forward of the model. This will interpolate the pre-trained
        position embeddings to the higher resolution.

    </Tip>
    c                   ó  ‡ — e Zd Zdeddfˆ fd„Z ee«       eee	e
e¬«      	 	 	 	 	 	 	 d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	f   fd„«       «       Zˆ xZS )ÚViTForImageClassificationr#   r%   Nc                 ó.  •— t         ‰| �  |«       |j                  | _        t        |d¬«      | _        |j                  dkD  r*t        j                  |j                  |j                  «      nt        j                  «       | _	        | j                  «        y )NF)r&  r   )r'   r(   Ú
num_labelsr%  r  r
   r¤   r,   ÚIdentityÚ
classifierr,  r¦   s     €r:   r(   z"ViTForImageClassification.__init__  ss   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ˜F°eÔ<ˆŒð OU×N_ÑN_ÐbcÒNcœ"Ÿ)™) F×$6Ñ$6¸×8IÑ8IÔJÔik×itÑitÓivˆŒð 	�‰Õr;   )r4  r5  r  r7  rY   r«   Úlabelsr¬   rû   rX   rü   c                 ób  — |�|n| j                   j                  }| j                  ||||||¬«      }|d   }	| j                  |	dd…ddd…f   «      }
d}|��¢|j	                  |
j
                  «      }| j                   j                  €�| j                  dk(  rd| j                   _        nl| j                  dkD  rL|j                  t        j                  k(  s|j                  t        j                  k(  rd| j                   _        nd| j                   _        | j                   j                  dk(  rIt        «       }| j                  dk(  r& ||
j                  «       |j                  «       «      }nŒ ||
|«      }n‚| j                   j                  dk(  r=t        «       } ||
j                  d| j                  «      |j                  d«      «      }n,| j                   j                  dk(  rt!        «       } ||
|«      }|s|
f|dd z   }|�|f|z   S |S t#        ||
|j$                  |j&                  ¬	«      S )
aŠ  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the image classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
        N)r«   r¬   rû   rX   rü   r   r   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationr@   )rZ  Úlogitsr´   r  )r#   r:  r  rj  r’   ÚdeviceÚproblem_typerh  rŒ   r*   Úlongro   r   Úsqueezer   rP   r   r   r´   r  )r8   rY   r«   rk  r¬   rû   rX   rü   r¼   r?  rp  rZ  Úloss_fctrÊ   s                 r:   rf   z!ViTForImageClassification.forward(  sî  € ð. &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—(‘(ØØØ/Ø!5Ø%=Ø#ð ó 
ˆð " !™*ˆà—‘ ²°A²q°Ñ!9Ó:ˆàˆØÑà—Y‘Y˜vŸ}™}Ó-ˆFØ�{‰{×'Ñ'Ð/Ø—?‘? aÒ'Ø/;�D—K‘KÕ,Ø—_‘_ qÒ(¨f¯l©l¼e¿j¹jÒ.HÈFÏLÉLÔ\a×\eÑ\eÒLeØ/L�D—K‘KÕ,à/K�D—K‘KÔ,à�{‰{×'Ñ'¨<Ò7Ü"›9�Ø—?‘? aÒ'Ù# F§N¡NÓ$4°f·n±nÓ6FÓG‘Dá# F¨FÓ3‘DØ—‘×)Ñ)Ð-JÒJÜ+Ó-�Ù §¡¨B°·±Ó @À&Ç+Á+ÈbÃ/ÓR‘Ø—‘×)Ñ)Ð-IÒIÜ,Ó.�Ù ¨Ó/�áØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä$ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r;   rB  )ri   rj   rk   r   r(   r   rC  r   Ú_IMAGE_CLASS_CHECKPOINTr   rE  Ú_IMAGE_CLASS_EXPECTED_OUTPUTr   r*   rn   rm   r	   r  rf   rq   rr   s   @r:   rf  rf    sï   ø„ ð 
˜yð 
¨Tõ 
ñ +Ð+?Ó@ÙØ*Ø)Ø$Ø4ô	ð 04Ø,0Ø)-Ø,0Ø/3Ø37Ø&*ñA
à˜uŸ|™|Ñ,ðA
ð ˜EŸL™LÑ)ðA
ð ˜Ÿ™Ñ&ð	A
ð
 $ D™>ðA
ð ' t™nðA
ð #+¨4¡.ðA
ð ˜d‘^ðA
ð 
ˆuÐ+Ð+Ñ	,òA
óó AôA
r;   rf  )rf  rO  r%  r  )r°   )Frl   Úcollections.abcry   r\  Útypingr   r   r   r   r   r   r	   r*   Útorch.utils.checkpointr
   Útorch.nnr   r   r   Úactivationsr   Úmodeling_outputsr   r   r   r   Úmodeling_utilsr   r   Úpytorch_utilsr   r   Úutilsr   r   r   r   r   r   Úconfiguration_vitr   Ú
get_loggerri   r²   rE  rD  rF  rv  rw  ÚModuler"   r0   rn   Úfloatr—   r™   r¾   rÇ   r×   rà   rä   ró   r  ÚVIT_START_DOCSTRINGrC  r%  r*  rO  rf  Ú__all__rþ   r;   r:   ú<module>r‡     s?  ðñ ã Û ß D× DÑ Dã Û Ý ß AÑ Aå !÷ó ÷ Gß Q÷÷ õ )ð 
ˆ×	Ñ	˜HÓ	%€ð €ð :Ð Ú&Ð ð 8Ð Ø-Ð ôU�B—I‘Iô Uôp$˜Ÿ™ô $ð\ ñ%Ø�I‰Ið%à�<‰<ð%ð 
�‰ð%ð �<‰<ð	%ð
 ˜UŸ\™\Ñ*ð%ð ð%ð ó%ô<;�r—y‘yô ;ô|�B—I‘Iô ô$$�2—9‘9ô $ôN�b—i‘iô ô �—	‘	ô ô'ˆr�y‰yô 'ôT0
�—‘ô 0
ôf)/˜ô )/ðX	Ð ðÐ ñ2 ØcØóô\
Ð!ó \
ó	ð\
ô~�—	‘	ô ñ ðð óôo
Ð 2ó o
óðo
ñd ðð óôU
Ð 2ó U
óðU
òp g�r;   