Ë
    S^(h;¦  ã                   óL  — d Z ddlZddlmZ ddl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%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	 dBd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>«      «       ZDe G d<„ d=e «      «       ZE e"d>e?«       G d?„ d@e>«      «       ZFg dA¢ZGy)CzPyTorch DeiT model.é    N)Ú	dataclass)ÚCallableÚ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)ÚModelOutputÚadd_code_sample_docstringsÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstringsÚ	torch_inté   )Ú
DeiTConfigr   z(facebook/deit-base-distilled-patch16-224)r   éÆ   i   ztabby, tabby 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 )ÚDeiTEmbeddingszv
    Construct the CLS token, distillation token, position and patch embeddings. Optionally, also the mask token.
    ÚconfigÚuse_mask_tokenÚreturnNc                 ó®  •— t         ‰| �  «        t        j                  t	        j
                  dd|j                  «      «      | _        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ÚzerosÚhidden_sizeÚ	cls_tokenÚdistillation_tokenÚ
mask_tokenÚDeiTPatchEmbeddingsÚpatch_embeddingsÚnum_patchesÚposition_embeddingsÚDropoutÚhidden_dropout_probÚdropoutÚ
patch_size)Úselfr#   r$   r3   Ú	__class__s       €úd/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/deit/modeling_deit.pyr)   zDeiTEmbeddings.__init__B   sç   ø€ Ü‰ÑÔäŸ™¤e§k¡k°!°Q¸×8JÑ8JÓ&KÓLˆŒÜ"$§,¡,¬u¯{©{¸1¸aÀ×ASÑASÓ/TÓ"UˆÔÙQ_œ"Ÿ,™,¤u§{¡{°1°a¸×9KÑ9KÓ'LÔMÐeiˆŒÜ 3°FÓ ;ˆÔØ×+Ñ+×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 and 2 class embeddings.

        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   r'   Néÿÿÿÿç      à?r   r   ÚbicubicF)ÚsizeÚmodeÚalign_corners©Údim)Úshaper4   r+   ÚjitÚ
is_tracingr8   r   ÚreshapeÚpermuter	   Ú
functionalÚinterpolateÚviewÚcat)r9   r=   r>   r?   r3   Únum_positionsÚclass_and_dist_pos_embedÚpatch_pos_embedrH   Ú
new_heightÚ	new_widthÚsqrt_num_positionss               r;   Úinterpolate_pos_encodingz'DeiTEmbeddings.interpolate_pos_encodingN   sb  € ð !×&Ñ& qÑ)¨AÑ-ˆØ×0Ñ0×6Ñ6°qÑ9¸AÑ=ˆô �y‰y×#Ñ#Ô%¨+¸Ò*FÈ6ÐUZÊ?Ø×+Ñ+Ð+à#'×#;Ñ#;ºA¸rÀ¸r¸EÑ#BÐ Ø×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Ð2°OÐDÈ!ÔLÐLr<   Úpixel_valuesÚbool_masked_posrX   c                 ó"  — |j                   \  }}}}| j                  |«      }|j                  «       \  }}	}|�K| j                  j	                  ||	d«      }
|j                  d«      j                  |
«      }|d|z
  z  |
|z  z   }| j                  j	                  |dd«      }| j                  j	                  |dd«      }t        j                  |||fd¬«      }| j                  }|r| j                  |||«      }||z   }| j                  |«      }|S )NrA   ç      ð?r   rG   )rI   r2   rD   r0   ÚexpandÚ	unsqueezeÚtype_asr.   r/   r+   rQ   r4   rX   r7   )r9   rY   rZ   rX   Ú_r>   r?   r=   Ú
batch_sizeÚ
seq_lengthÚmask_tokensÚmaskÚ
cls_tokensÚdistillation_tokensÚposition_embeddings                  r;   ÚforwardzDeiTEmbeddings.forwardv   s  € ð +×0Ñ0Ñˆˆ1ˆf�eØ×*Ñ*¨<Ó8ˆ
à$.§O¡OÓ$5Ñ!ˆ
�J àÐ&ØŸ/™/×0Ñ0°¸ZÈÓLˆKà"×,Ñ,¨RÓ0×8Ñ8¸ÓEˆDØ# s¨T¡zÑ2°[À4Ñ5GÑGˆJà—^‘^×*Ñ*¨:°r¸2Ó>ˆ
à"×5Ñ5×<Ñ<¸ZÈÈRÓPÐä—Y‘Y 
Ð,?ÀÐLÐRSÔTˆ
Ø!×5Ñ5Ðá#Ø!%×!>Ñ!>¸zÈ6ÐSXÓ!YÐàÐ"4Ñ4ˆ
Ø—\‘\ *Ó-ˆ
ØÐr<   )F©NF)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úboolr)   r+   ÚTensorÚintrX   r   Ú
BoolTensorrh   Ú__classcell__©r:   s   @r;   r"   r"   =   s›   ø„ ññ
,˜zð 
,¸4ð 
,ÈDõ 
,ð&M°5·<±<ð &MÈð &MÐUXð &MÐ]b×]iÑ]ió &MðV 7;Ø).ñ	à—l‘lðð " %×"2Ñ"2Ñ3ðð #'ð	ð
 
�‰÷r<   r"   c                   óZ   ‡ — e Zd ZdZˆ fd„Zdej                  dej                  fd„Zˆ xZS )r1   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_sizer8   Únum_channelsr-   Ú
isinstanceÚcollectionsÚabcÚIterabler3   r	   ÚConv2dÚ
projection)r9   r#   rx   r8   ry   r-   r3   r:   s          €r;   r)   zDeiTPatchEmbeddings.__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   r%   c                 ó¼   — |j                   \  }}}}|| j                  k7  rt        d«      ‚| j                  |«      j	                  d«      j                  dd«      }|S )NzeMake sure that the channel dimension of the pixel values match with the one set in the configuration.r'   r   )rI   ry   Ú
ValueErrorr   ÚflattenÚ	transpose)r9   rY   ra   ry   r>   r?   Úxs          r;   rh   zDeiTPatchEmbeddings.forward¬   sa   € Ø2>×2DÑ2DÑ/ˆ
�L &¨%Ø˜4×,Ñ,Ò,ÜØwóð ð �O‰O˜LÓ)×1Ñ1°!Ó4×>Ñ>¸qÀ!ÓDˆØˆr<   )	rj   rk   rl   rm   r)   r+   ro   rh   rr   rs   s   @r;   r1   r1   –   s)   ø„ ñôjð E§L¡Lð °U·\±\÷ r<   r1   ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingr7   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 )NrA   éþÿÿÿ)rH   Údtype)ÚpÚtrainingr   r'   )r+   Úmatmulrƒ   r	   rN   ÚsoftmaxÚfloat32Útor�   r7   r�   Ú
contiguous)
r…   r†   r‡   rˆ   r‰   rŠ   r7   Ú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 )ÚDeiTSelfAttentionr#   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 ú.g      à¿F)Úbias)r(   r)   r-   Únum_attention_headsÚhasattrr�   r#   rp   Úattention_head_sizeÚall_head_sizeÚattention_probs_dropout_probÚdropout_probrŠ   Ú	is_causalr	   ÚLinearÚqkv_biasr†   r‡   rˆ   ©r9   r#   r:   s     €r;   r)   zDeiTSelfAttention.__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<   r„   c                 ó¤   — |j                  «       d d | j                  | j                  fz   }|j                  |«      }|j	                  dddd«      S )NrA   r   r'   r   r   )rD   rŸ   r¡   rP   rM   )r9   r„   Únew_x_shapes      r;   Útranspose_for_scoresz&DeiTSelfAttention.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Š   r7   rŒ   )r«   r‡   rˆ   r†   r˜   r#   Ú_attn_implementationÚloggerÚwarning_oncer   r¥   rŠ   r�   r¤   rD   r¢   rL   )r9   Úhidden_statesr¬   r­   Ú	key_layerÚvalue_layerÚquery_layerÚattention_interfaceÚcontext_layerÚattention_probsÚnew_context_layer_shapeÚoutputss               r;   rh   zDeiTSelfAttention.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<   ri   )rj   rk   rl   r   r)   r+   ro   r«   r   rn   r   r   rh   rr   rs   s   @r;   rš   rš   Ö   s†   ø„ ð]˜zð ]¨dõ ]ð(% 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 )	ÚDeiTSelfOutputz¡
    The residual connection is defined in DeiTLayer 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-   Údenser5   r6   r7   r¨   s     €r;   r)   zDeiTSelfOutput.__init__  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r<   rµ   Úinput_tensorc                 óJ   — | j                  |«      }| j                  |«      }|S rÁ   ©rÂ   r7   ©r9   rµ   rÃ   s      r;   rh   zDeiTSelfOutput.forward   s$   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆàÐr<   )
rj   rk   rl   rm   r   r)   r+   ro   rh   rr   rs   s   @r;   r¿   r¿     sD   ø„ ñð
>˜zð >¨dõ >ð
 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 )ÚDeiTAttentionr#   r%   Nc                 ó€   •— t         ‰| �  «        t        |«      | _        t	        |«      | _        t        «       | _        y rÁ   )r(   r)   rš   Ú	attentionr¿   ÚoutputÚsetÚpruned_headsr¨   s     €r;   r)   zDeiTAttention.__init__)  s0   ø€ Ü‰ÑÔÜ*¨6Ó2ˆŒÜ$ VÓ,ˆŒÜ›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)r9   rÎ   Úindexs      r;   Úprune_headszDeiTAttention.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<   rµ   r¬   r­   c                 óh   — | j                  |||«      }| j                  |d   |«      }|f|dd  z   }|S )Nr   r   )rÊ   rË   )r9   rµ   r¬   r­   Úself_outputsÚattention_outputr½   s          r;   rh   zDeiTAttention.forwardA  sE   € ð —~‘~ m°YÐ@QÓRˆàŸ;™; |°A¡¸ÓFÐà#Ð%¨°Q°RÐ(8Ñ8ˆØˆr<   ri   )rj   rk   rl   r   r)   r   rp   rÓ   r+   ro   r   rn   r   r   rh   rr   rs   s   @r;   rÈ   rÈ   (  s’   ø„ ð"˜zð "¨dõ "ð;  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 )ÚDeiTIntermediater#   r%   Nc                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y rÁ   )r(   r)   r	   r¦   r-   Úintermediate_sizerÂ   rz   Ú
hidden_actÚstrr   Úintermediate_act_fnr¨   s     €r;   r)   zDeiTIntermediate.__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Ý   )r9   rµ   s     r;   rh   zDeiTIntermediate.forwardY  s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆàÐr<   ©	rj   rk   rl   r   r)   r+   ro   rh   rr   rs   s   @r;   rØ   rØ   P  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 )Ú
DeiTOutputr#   r%   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j                  «      | _	        y rÁ   )
r(   r)   r	   r¦   rÚ   r-   rÂ   r5   r6   r7   r¨   s     €r;   r)   zDeiTOutput.__init__b  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;   rh   zDeiTOutput.forwardg  s.   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆà%¨Ñ4ˆàÐr<   rß   rs   s   @r;   rá   rá   a  s?   ø„ ð>˜zð >¨dõ >ð
 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 )Ú	DeiTLayerz?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DeiTLayer.__init__t  s‡   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ& vÓ.ˆŒÜ,¨VÓ4ˆÔÜ  Ó(ˆŒÜ "§¡¨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Ë   )r9   rµ   r¬   r­   Úself_attention_outputsrÖ   r½   Úlayer_outputs           r;   rh   zDeiTLayer.forward~  s–   € ð "&§¡Ø×!Ñ! -Ó0ØØ/ð "0ó "
Ðð
 2°!Ñ4ÐØ(¨¨Ð,ˆð )¨=Ñ8ˆð ×+Ñ+¨MÓ:ˆØ×(Ñ(¨Ó6ˆð —{‘{ <°Ó?ˆà�/ GÑ+ˆàˆr<   ri   )rj   rk   rl   rm   r   r)   r+   ro   r   rn   r   r   rh   rr   rs   s   @r;   rå   rå   q  s�   ø„ ÙIð[˜zð [¨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dej                  deej                     deded	ede	e
ef   fd
„Zˆ xZS )ÚDeiTEncoderr#   r%   Nc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w ri   )
r(   r)   r#   r	   Ú
ModuleListÚrangeÚnum_hidden_layersrå   ÚlayerÚgradient_checkpointing)r9   r#   r`   r:   s      €r;   r)   zDeiTEncoder.__init__�  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]¼uÀV×E]ÑE]Ó?^Ö#_¸!¤I¨fÕ$5Ò#_Ó`ˆŒ
Ø&+ˆÕ#ùò $`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&DeiTEncoder.forward.<locals>.<genexpr>Ç  s   è ø€ Òm˜qÐ_`Ñ_lœÑmùs   ‚Š)Úlast_hidden_staterµ   Ú
attentions)Ú	enumeraterù   rú   r�   Ú_gradient_checkpointing_funcÚ__call__Útupler   )r9   rµ   r¬   r­   rû   rü   Úall_hidden_statesÚall_self_attentionsÚiÚlayer_moduleÚlayer_head_maskÚlayer_outputss               r;   rh   zDeiTEncoder.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)rj   rk   rl   r   r)   r+   ro   r   rn   r   r  r   rh   rr   rs   s   @r;   rô   rô   œ  sz   ø„ ð,˜zð ,¨dõ ,ð -1Ø"'Ø%*Ø ñ)
à—|‘|ð)
ð ˜EŸL™LÑ)ð)
ð  ð	)
ð
 #ð)
ð ð)
ð 
ˆu�oÐ%Ñ	&÷)
r<   rô   c                   ó†   — e Zd ZdZeZdZdZdZdgZ	dZ
dZdeej                  ej                  ej                   f   ddfd	„Zy)
ÚDeiTPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚdeitrY   Trå   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$        «      rž|j&                  j                  j                  «        |j(                  j                  j                  «        |j*                  j                  j                  «        |j,                  �%|j,                  j                  j                  «        yyy)zInitialize the weightsr±   )ÚmeanÚstdNr\   )rz   r	   r¦   r~   ÚinitÚtrunc_normal_ÚweightÚdatar“   r+   r’   r#   Úinitializer_ranger�   rž   Úzero_rì   Úfill_r"   r.   r4   r/   r0   )r9   r…   s     r;   Ú_init_weightsz!DeiTPreTrainedModel._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Õ)Ü˜¤Ô/Ø×Ñ×!Ñ!×'Ñ'Ô)Ø×&Ñ&×+Ñ+×1Ñ1Ô3Ø×%Ñ%×*Ñ*×0Ñ0Ô2Ø× Ñ Ð,Ø×!Ñ!×&Ñ&×,Ñ,Õ.ð -ð	 0r<   )rj   rk   rl   rm   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  Ï  s^   „ ñð
 €LØÐØ$€OØ&*Ð#Ø$˜ÐØ€NØ!Ðð/ E¨"¯)©)°R·Y±YÀÇÁÐ*LÑ$Mð /ÐRVô /r<   r  aF  
    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 ([`DeiTConfig`]): 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
            [`DeiTImageProcessor.__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.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
        interpolate_pos_encoding (`bool`, *optional*, defaults to `False`):
            Whether to interpolate the pre-trained position encodings.
z^The bare DeiT Model transformer outputting raw hidden-states without any specific head on top.c                   ó  ‡ — e Zd Zddedededdfˆ fd„Zdefd„Z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deeef   fd„«       «       Zˆ xZS )Ú	DeiTModelr#   Úadd_pooling_layerr$   r%   Nc                 ó  •— 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Ú
DeiTPoolerÚpoolerÚ	post_init)r9   r#   r&  r$   r:   s       €r;   r)   zDeiTModel.__init__  sk   ø€ Ü‰Ñ˜Ô ØˆŒä(¨ÀÔOˆŒÜ" 6Ó*ˆŒäŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÙ,=”j Ô(À4ˆŒð 	�‰Õr<   c                 ó.   — | j                   j                  S rÁ   )r=   r2   )r9   s    r;   Úget_input_embeddingszDeiTModel.get_input_embeddings(  s   € Ø�‰×/Ñ/Ð/r<   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(  rù   rÊ   rÓ   )r9   Úheads_to_prunerù   rÎ   s       r;   Ú_prune_headszDeiTModel._prune_heads+  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û   rü   rX   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=   r2   r   r  r�   r“   r(  r)  r+  r   rµ   r  )r9   rY   rZ   r¬   r­   rû   rü   rX   Úexpected_dtypeÚembedding_outputÚencoder_outputsÚsequence_outputÚpooled_outputÚhead_outputss                 r;   rh   zDeiTModel.forward3  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©NNNNNNF)rj   rk   rl   r   rn   r)   r1   r.  r2  r   ÚDEIT_INPUTS_DOCSTRINGr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCÚ_EXPECTED_OUTPUT_SHAPEr   r+   ro   rq   r   r   rh   rr   rs   s   @r;   r%  r%    s	  ø„ ñ
˜zð ¸dð Ð[_ð Ðlpõ ð0Ð&9ó 0òCñ +Ð+@ÓAÙØ&Ø.Ø$ØØ.ôð 04Ø6:Ø,0Ø,0Ø/3Ø&*Ø).ñ;
à˜uŸ|™|Ñ,ð;
ð " %×"2Ñ"2Ñ3ð;
ð ˜EŸL™LÑ)ð	;
ð
 $ D™>ð;
ð ' t™nð;
ð ˜d‘^ð;
ð #'ð;
ð 
ˆuÐ0Ð0Ñ	1ò;
óó Bô;
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DeiTPooler.__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  )r9   rµ   Úfirst_token_tensorr@  s       r;   rh   zDeiTPooler.forward€  s6   € ð +ª1¨a¨4Ñ0ÐØŸ
™
Ð#5Ó6ˆØŸ™¨Ó6ˆØÐr<   )rj   rk   rl   r   r)   rh   rr   rs   s   @r;   r*  r*  z  s   ø„ ð4˜zõ 4ö
r<   r*  aW  DeiT 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deeef   fd„«       «       Zˆ xZS )ÚDeiTForMaskedImageModelingr#   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$   r'   r   )Úin_channelsÚout_channelsrv   )r(   r)   r%  r  r	   Ú
Sequentialr~   r-   Úencoder_stridery   ÚPixelShuffleÚdecoderr,  r¨   s     €r;   r)   z#DeiTForMaskedImageModeling.__init__–  s�   ø€ Ü‰Ñ˜Ô ä˜f¸ÈdÔSˆŒ	ä—}‘}Ü�I‰IØ"×.Ñ.Ø#×2Ñ2°AÑ5¸×8KÑ8KÑKØôô
 �O‰O˜F×1Ñ1Ó2ó
ˆŒð 	�‰Õr<   ©r5  r  rY   rZ   r¬   r­   rû   rü   rX   c           	      óº  — |�|n| j                   j                  }| j                  |||||||¬«      }|d   }	|	dd…dd…f   }	|	j                  \  }
}}t	        |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 )aM  
        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, DeiTForMaskedImageModeling
        >>> 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("facebook/deit-base-distilled-patch16-224")
        >>> model = DeiTForMaskedImageModeling.from_pretrained("facebook/deit-base-distilled-patch16-224")

        >>> 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]
        ```N)rZ   r¬   r­   rû   rü   rX   r   r   rA   rB   r'   Únone)Ú	reductiongñhãˆµøä>)ÚlossÚreconstructionrµ   r  )r#   r:  r  rI   rp   rM   rL   rV  rx   r8   Úrepeat_interleaver^   r”   r	   rN   Úl1_lossÚsumry   r   rµ   r  )r9   rY   rZ   r¬   r­   rû   rü   rX   r½   r?  ra   Úsequence_lengthry   r>   r?   Úreconstructed_pixel_valuesÚmasked_im_lossrD   rd   Úreconstruction_lossrË   s                        r;   rh   z"DeiTForMaskedImageModeling.forward§  sí  € ðR &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ+ØØ/Ø!5Ø#Ø%=ð ó 
ˆð " !™*ˆð *ª!¨Q¨r¨T¨'Ñ2ˆØ4C×4IÑ4IÑ1ˆ
�O \Ü˜_¨cÑ1Ó2Ð2ˆ�Ø)×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  )rj   rk   rl   r   r)   r   rC  r   r   rE  r   r+   ro   rq   rn   r   r  rh   rr   rs   s   @r;   rO  rO  ‰  sæ   ø„ ð˜zð ¨dõ ñ" +Ð+@ÓAÙÐ+DÐSbÔcð 04Ø6:Ø,0Ø,0Ø/3Ø&*Ø).ñT
à˜uŸ|™|Ñ,ðT
ð " %×"2Ñ"2Ñ3ðT
ð ˜EŸL™LÑ)ð	T
ð
 $ D™>ðT
ð ' t™nðT
ð ˜d‘^ðT
ð #'ðT
ð 
ˆuÐ/Ð/Ñ	0òT
ó dó BôT
r<   rO  z¥
    DeiT 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.
    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deeef   fd„«       «       Zˆ xZS )ÚDeiTForImageClassificationr#   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#DeiTForImageClassification.__init__  ss   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ˜f¸Ô>ˆŒ	ð OU×N_ÑN_ÐbcÒNcœ"Ÿ)™) F×$6Ñ$6¸×8IÑ8IÔJÔik×itÑitÓivˆŒð 	�‰Õr<   rW  rY   r¬   Úlabelsr­   rû   rü   rX   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 )
al  
        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).

        Returns:

        Examples:

        ```python
        >>> from transformers import AutoImageProcessor, DeiTForImageClassification
        >>> import torch
        >>> from PIL import Image
        >>> import requests

        >>> torch.manual_seed(3)  # doctest: +IGNORE_RESULT
        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)

        >>> # note: we are loading a DeiTForImageClassificationWithTeacher from the hub here,
        >>> # so the head will be randomly initialized, hence the predictions will be random
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/deit-base-distilled-patch16-224")
        >>> model = DeiTForImageClassification.from_pretrained("facebook/deit-base-distilled-patch16-224")

        >>> inputs = image_processor(images=image, return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> logits = outputs.logits
        >>> # model predicts one of the 1000 ImageNet classes
        >>> predicted_class_idx = logits.argmax(-1).item()
        >>> print("Predicted class:", model.config.id2label[predicted_class_idx])
        Predicted class: Polaroid camera, Polaroid Land camera
        ```N©r¬   r­   rû   rü   rX   r   r   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationrA   )r[  Úlogitsrµ   r  )r#   r:  r  rj  r“   ÚdeviceÚproblem_typerh  r�   r+   Úlongrp   r   Úsqueezer   rP   r
   r   rµ   r  )r9   rY   r¬   rk  r­   rû   rü   rX   r½   r?  rq  r[  Úloss_fctrË   s                 r;   rh   z"DeiTForImageClassification.forward  sñ  € ðZ &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  )rj   rk   rl   r   r)   r   rC  r   r   rE  r   r+   ro   rn   r   r  rh   rr   rs   s   @r;   re  re     sã   ø„ ð
˜zð 
¨dõ 
ñ +Ð+@ÓAÙÐ+@ÈÔ_ð 04Ø,0Ø)-Ø,0Ø/3Ø&*Ø).ñ[
à˜uŸ|™|Ñ,ð[
ð ˜EŸL™LÑ)ð[
ð ˜Ÿ™Ñ&ð	[
ð
 $ D™>ð[
ð ' t™nð[
ð ˜d‘^ð[
ð #'ð[
ð 
ˆuÐ+Ð+Ñ	,ò[
ó `ó Bô[
r<   re  c                   óæ   — 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j                     ed<   dZeeej                        ed<   dZeeej                        ed<   y)Ú+DeiTForImageClassificationWithTeacherOutputa5  
    Output type of [`DeiTForImageClassificationWithTeacher`].

    Args:
        logits (`torch.FloatTensor` of shape `(batch_size, config.num_labels)`):
            Prediction scores as the average of the cls_logits and distillation logits.
        cls_logits (`torch.FloatTensor` of shape `(batch_size, config.num_labels)`):
            Prediction scores of the classification head (i.e. the linear layer on top of the final hidden state of the
            class token).
        distillation_logits (`torch.FloatTensor` of shape `(batch_size, config.num_labels)`):
            Prediction scores of the distillation head (i.e. the linear layer on top of the final hidden state of the
            distillation token).
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
            shape `(batch_size, sequence_length, hidden_size)`. Hidden-states of the model at the output of each layer
            plus the initial embedding outputs.
        attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
            sequence_length)`. Attentions weights after the attention softmax, used to compute the weighted average in
            the self-attention heads.
    Nrq  Ú
cls_logitsÚdistillation_logitsrµ   r  )rj   rk   rl   rm   rq  r   r+   ÚFloatTensorÚ__annotations__ry  rz  rµ   r   r  rþ   r<   r;   rx  rx  t  s}   … ñð, +/€FˆH�U×&Ñ&Ñ'Ó.Ø.2€J�˜×*Ñ*Ñ+Ó2Ø7;Ð˜ %×"3Ñ"3Ñ4Ó;Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r<   rx  aˆ  
    DeiT Model transformer with image classification heads on top (a linear layer on top of the final hidden state of
    the [CLS] token and a linear layer on top of the final hidden state of the distillation token) e.g. for ImageNet.

    .. warning::

           This model supports inference-only. Fine-tuning with distillation (i.e. with a teacher) is not yet
           supported.
    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   d	ee   d
ee   dedeee	f   fd„«       «       Zˆ xZS )Ú%DeiTForImageClassificationWithTeacherr#   r%   Nc                 óÒ  •— t         ‰| �  |«       |j                  | _        t        |d¬«      | _        |j                  dkD  r*t        j                  |j                  |j                  «      nt        j                  «       | _	        |j                  dkD  r*t        j                  |j                  |j                  «      nt        j                  «       | _
        | j                  «        y rg  )r(   r)   rh  r%  r  r	   r¦   r-   ri  Úcls_classifierÚdistillation_classifierr,  r¨   s     €r;   r)   z.DeiTForImageClassificationWithTeacher.__init__   s·   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ˜f¸Ô>ˆŒ	ð AG×@QÑ@QÐTUÒ@UŒB�I‰I�f×(Ñ(¨&×*;Ñ*;Ô<Ô[]×[fÑ[fÓ[hð 	Ôð AG×@QÑ@QÐTUÒ@UŒB�I‰I�f×(Ñ(¨&×*;Ñ*;Ô<Ô[]×[fÑ[fÓ[hð 	Ô$ð
 	�‰Õr<   )r4  r5  r  r7  rY   r¬   r­   rû   rü   rX   c                 óP  — |�|n| j                   j                  }| j                  ||||||¬«      }|d   }| j                  |d d …dd d …f   «      }	| j	                  |d d …dd d …f   «      }
|	|
z   dz  }|s||	|
f|dd  z   }|S t        ||	|
|j                  |j                  ¬«      S )Nrm  r   r   r'   )rq  ry  rz  rµ   r  )r#   r:  r  r€  r�  rx  rµ   r  )r9   rY   r¬   r­   rû   rü   rX   r½   r?  ry  rz  rq  rË   s                r;   rh   z-DeiTForImageClassificationWithTeacher.forward±  sÜ   € ð  &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØØ/Ø!5Ø#Ø%=ð ó 
ˆð " !™*ˆà×(Ñ(¨º¸Aºq¸Ñ)AÓBˆ
Ø"×:Ñ:¸?Ê1ÈaÒQRÈ7Ñ;SÓTÐð Ð2Ñ2°aÑ7ˆáØ˜jÐ*=Ð>ÀÈÈÀÑLˆFØˆMä:ØØ!Ø 3Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r<   )NNNNNF)rj   rk   rl   r   r)   r   rC  r   Ú_IMAGE_CLASS_CHECKPOINTrx  rE  Ú_IMAGE_CLASS_EXPECTED_OUTPUTr   r+   ro   rn   r   r  rh   rr   rs   s   @r;   r~  r~  “  sË   ø„ ð˜zð ¨dõ ñ" +Ð+@ÓAÙØ*Ø?Ø$Ø4ô	ð 04Ø,0Ø,0Ø/3Ø&*Ø).ñ&
à˜uŸ|™|Ñ,ð&
ð ˜EŸL™LÑ)ð&
ð $ D™>ð	&
ð
 ' t™nð&
ð ˜d‘^ð&
ð #'ð&
ð 
ˆuÐAÐAÑ	Bò&
óó Bô&
r<   r~  )re  r~  rO  r%  r  )r±   )Hrm   Úcollections.abcr{   Údataclassesr   Útypingr   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   r   Úconfiguration_deitr   Ú
get_loggerrj   r³   rE  rD  rF  rƒ  r„  ÚModuler"   r1   ro   Úfloatr˜   rš   r¿   rÈ   rØ   rá   rå   rô   r  ÚDEIT_START_DOCSTRINGrC  r%  r*  rO  re  rx  r~  Ú__all__rþ   r<   r;   ú<module>r•     s‰  ðñ ã Ý !ß 8Õ 8ã Û Ý ß AÑ Aå !÷ó ÷ Gß Q÷÷ ñ õ +ð 
ˆ×	Ñ	˜HÓ	%€ð €ð AÐ Ú&Ð ð EÐ Ø1Ð ôV�R—Y‘Yô Vôr˜"Ÿ)™)ô ðP ñ%Ø�I‰Ið%à�<‰<ð%ð 
�‰ð%ð �<‰<ð	%ð
 ˜UŸ\™\Ñ*ð%ð ð%ð ó%ô>;˜Ÿ	™	ô ;ô~�R—Y‘Yô ô&$�B—I‘Iô $ôP�r—y‘yô ô"�—‘ô ô '�—	‘	ô 'ôV0
�"—)‘)ô 0
ôf /˜/ô  /ðF	Ð ðÐ ñ2 ØdØóô\
Ð#ó \
ó	ð\
ô@�—‘ô ñ ðð óôh
Ð!4ó h
óðh
ñV ðð óôj
Ð!4ó j
óðj
ðZ ô:°+ó :ó ð:ñ< ðð óô?
Ð,?ó ?
óð?
òD�r<   