Ë
    T^(h…y  ã                   ó´  — 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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 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! d
dl"m#Z#  e jH                  e%«      Z&dZ'dZ( G d„ dejR                  «      Z* G d„ dejR                  «      Z+ G d„ de«      Z,	 d7dejR                  de
jZ                  de
jZ                  de
jZ                  dee
jZ                     de.de.fd„Z/ G d„ dejR                  «      Z0 G d„ dejR                  «      Z1 G d „ d!ejR                  «      Z2 G d"„ d#ejR                  «      Z3 G d$„ d%ejR                  «      Z4 G d&„ d'ejR                  «      Z5 G d(„ d)ejR                  «      Z6 G d*„ d+ejR                  «      Z7d,Z8g d-¢Z9d.Z: ed/e:«       G d0„ d1e,«      «       Z;dZ<d2Z= ed3e:«       G d4„ d5e,«      «       Z>g d6¢Z?y)8é    N)ÚCallableÚDictÚListÚOptionalÚSetÚTupleÚUnion)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )ÚACT2FN)ÚBaseModelOutputÚBaseModelOutputWithPoolingÚImageClassifierOutput)Ú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Ú	torch_inté   )ÚIJepaConfigzfacebook/ijepa_vith14_1kr   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 )ÚIJepaPatchEmbeddingszì
    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)ÚsuperÚ__init__Ú
image_sizeÚ
patch_sizeÚnum_channelsÚhidden_sizeÚ
isinstanceÚcollectionsÚabcÚIterableÚnum_patchesÚnnÚConv2dÚ
projection)ÚselfÚconfigr$   r%   r&   r'   r,   Ú	__class__s          €úf/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/ijepa/modeling_ijepa.pyr#   zIJepaPatchEmbeddings.__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ˆ�ó    Úpixel_valuesÚinterpolate_pos_encodingÚreturnc                 óŽ  — |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).é   )Úshaper&   Ú
ValueErrorr$   r/   ÚflattenÚ	transpose)r0   r5   r6   Ú
batch_sizer&   ÚheightÚwidthÚ
embeddingss           r3   ÚforwardzIJepaPatchEmbeddings.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ˆ
ØÐr4   ©F)
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r#   ÚtorchÚTensorÚboolrD   Ú__classcell__©r2   s   @r3   r   r   %   s3   ø„ ñôjñ E§L¡Lð ÈDð Ð]b×]iÑ]i÷ r4   r   c            	       óÒ   ‡ — e Zd ZdZddededdfˆ fd„Zdej                  de	d	e	dej                  fd
„Z
	 	 ddej                  deej                     dedej                  fd„Zˆ xZS )ÚIJepaEmbeddingszb
    Construct the CLS token, position and patch embeddings. Optionally, also the mask token.
    r1   Úuse_mask_tokenr7   Nc                 óÒ  •— t         ‰| �  «        |r4t        j                  t	        j
                  dd|j                  «      «      nd | _        t        |«      | _	        | j                  j                  }t        j                  t	        j                  d||j                  «      «      | _        t        j                  |j                  «      | _        |j                   | _        || _        y )Nr   )r"   r#   r-   Ú	ParameterrJ   Úzerosr'   Ú
mask_tokenr   Úpatch_embeddingsr,   ÚrandnÚposition_embeddingsÚDropoutÚhidden_dropout_probÚdropoutr%   r1   )r0   r1   rQ   r,   r2   s       €r3   r#   zIJepaEmbeddings.__init__Q   s¢   ø€ Ü‰ÑÔÙQ_œ"Ÿ,™,¤u§{¡{°1°a¸×9KÑ9KÓ'LÔMÐeiˆŒÜ 4°VÓ <ˆÔØ×+Ñ+×7Ñ7ˆÜ#%§<¡<´·±¸A¸{ÈF×L^ÑL^Ó0_Ó#`ˆÔ Ü—z‘z &×"<Ñ"<Ó=ˆŒØ ×+Ñ+ˆŒØˆ�r4   rC   rA   rB   c                 ó0  — |j                   d   }| j                  j                   d   }t        j                  j	                  «       s||k(  r||k(  r| j                  S | j                  }|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|«      }|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   éÿÿÿÿg      à?r   r   r;   ÚbicubicF)ÚsizeÚmodeÚalign_corners)r<   rX   rJ   ÚjitÚ
is_tracingr%   r   ÚreshapeÚpermuter-   Ú
functionalÚinterpolateÚview)r0   rC   rA   rB   r,   Únum_positionsÚpatch_pos_embedÚdimÚ
new_heightÚ	new_widthÚsqrt_num_positionss              r3   r6   z(IJepaEmbeddings.interpolate_pos_encoding[   s#  € ð !×&Ñ& qÑ)ˆØ×0Ñ0×6Ñ6°qÑ9ˆô �y‰y×#Ñ#Ô%¨+¸Ò*FÈ6ÐUZÊ?Ø×+Ñ+Ð+à×2Ñ2ˆà×Ñ˜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ˆàÐr4   r5   Úbool_masked_posr6   c                 óx  — |j                   \  }}}}| j                  ||¬«      }|�Z|j                   d   }	| j                  j                  ||	d«      }
|j	                  d«      j                  |
«      }|d|z
  z  |
|z  z   }|r|| j                  |||«      z   }n|| j                  z   }| j                  |«      }|S )N)r6   r   r]   ç      ð?)	r<   rV   rU   ÚexpandÚ	unsqueezeÚtype_asr6   rX   r[   )r0   r5   ro   r6   r@   Ú_rA   rB   rC   Ú
seq_lengthÚmask_tokensÚmasks               r3   rD   zIJepaEmbeddings.forward‚   sÓ   € ð (4×'9Ñ'9Ñ$ˆ
�A�v˜uØ×*Ñ*¨<ÐRjÐ*Ókˆ
àÐ&Ø#×)Ñ)¨!Ñ,ˆJØŸ/™/×0Ñ0°¸ZÈÓLˆKà"×,Ñ,¨RÓ0×8Ñ8¸ÓEˆDØ# s¨T¡zÑ2°[À4Ñ5GÑGˆJñ $Ø# d×&CÑ&CÀJÐPVÐX]Ó&^Ñ^‰Jà# d×&>Ñ&>Ñ>ˆJà—\‘\ *Ó-ˆ
àÐr4   rE   ©NF)rF   rG   rH   rI   r   rL   r#   rJ   rK   Úintr6   r   Ú
BoolTensorrD   rM   rN   s   @r3   rP   rP   L   s–   ø„ ññ˜{ð ¸Dð ÈTõ ð%°5·<±<ð %Èð %ÐUXð %Ð]b×]iÑ]ió %ðT 7;Ø).ñ	à—l‘lðð " %×"2Ñ"2Ñ3ðð #'ð	ð
 
�‰÷r4   rP   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	)ÚIJepaPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úijepar5   TrP   Ú
IJepaLayerÚmoduler7   Nc                 ól  — 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Ë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 weightsç        )ÚmeanÚstdNrq   )r(   r-   ÚLinearr.   ÚinitÚtrunc_normal_ÚweightÚdataÚtorJ   Úfloat32r1   Úinitializer_rangeÚdtypeÚbiasÚzero_Ú	LayerNormÚfill_rP   rX   rU   )r0   r€   s     r3   Ú_init_weightsz"IJepaPreTrainedModel._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Ü.0¯g©g×.CÑ.CØ×*Ñ*×/Ñ/×2Ñ2´5·=±=ÓAØØ—K‘K×1Ñ1ð /Dó /÷ ‰b�×+Ñ+×1Ñ1Ó2ð	 ×&Ñ&Ô+ð
 × Ñ Ð,Ø×!Ñ!×&Ñ&×,Ñ,Õ.ð -ð 1r4   )rF   rG   rH   rI   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’   © r4   r3   r}   r}   �   sa   „ ñð
 €LØÐØ$€OØ&*Ð#Ø*¨LÐ9ÐØ€NØ!Ðð/ E¨"¯)©)°R·Y±YÀÇÁÐ*LÑ$Mð /ÐRVô /r4   r}   r€   ÚqueryÚkeyÚvalueÚattention_maskÚscalingr[   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]   éþÿÿÿ)rk   r�   )ÚpÚtrainingr   r;   )rJ   Úmatmulr?   r-   rf   Úsoftmaxr‹   rŠ   r�   r[   r£   Ú
contiguous)
r€   r›   rœ   r�   rž   rŸ   r[   ÚkwargsÚattn_weightsÚattn_outputs
             r3   Ú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à˜Ð$Ð$r4   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 )ÚIJepaSelfAttentionr1   r7   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 r9   g      à¿F)rŽ   )r"   r#   r'   Únum_attention_headsÚhasattrr=   r1   rz   Úattention_head_sizeÚall_head_sizeÚattention_probs_dropout_probÚdropout_probrŸ   Ú	is_causalr-   r…   Úqkv_biasr›   rœ   r�   ©r0   r1   r2   s     €r3   r#   zIJepaSelfAttention.__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Ô\ˆ�
r4   Úxc                 ó¤   — |j                  «       d d | j                  | j                  fz   }|j                  |«      }|j	                  dddd«      S )Nr]   r   r;   r   r   )r_   r¯   r±   rh   re   )r0   r¸   Únew_x_shapes      r3   Útranspose_for_scoresz'IJepaSelfAttention.transpose_for_scoresõ   sL   € Ø—f‘f“h˜s �m t×'?Ñ'?À×AYÑAYÐ&ZÑZˆØ�F‰F�;ÓˆØ�y‰y˜˜A˜q !Ó$Ð$r4   Ú	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µ   rŸ   r[   r¡   )r»   rœ   r�   r›   rª   r1   Ú_attn_implementationÚloggerÚwarning_oncer   rµ   rŸ   r£   r´   r_   r²   rd   )r0   Úhidden_statesr¼   r½   Ú	key_layerÚvalue_layerÚquery_layerÚattention_interfaceÚcontext_layerÚattention_probsÚnew_context_layer_shapeÚoutputss               r3   rD   zIJepaSelfAttention.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]ˆàˆr4   ry   )rF   rG   rH   r   r#   rJ   rK   r»   r   rL   r	   r   rD   rM   rN   s   @r3   r¬   r¬   à   s†   ø„ ð]˜{ð ]¨tõ ]ð(% e§l¡lð %°u·|±|ó %ð bgñ!Ø(0°·±Ñ(>ð!ØZ^ð!à	ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷!r4   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 )	ÚIJepaSelfOutputz¢
    The residual connection is defined in IJepaLayer instead of here (as is the case with other models), due to the
    layernorm applied before each block.
    r1   r7   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  |j                  «      | _        y ©N)	r"   r#   r-   r…   r'   ÚdenserY   rZ   r[   r·   s     €r3   r#   zIJepaSelfOutput.__init__$  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r4   rÄ   Úinput_tensorc                 óJ   — | j                  |«      }| j                  |«      }|S rÐ   ©rÑ   r[   ©r0   rÄ   rÒ   s      r3   rD   zIJepaSelfOutput.forward)  s$   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆàÐr4   )
rF   rG   rH   rI   r   r#   rJ   rK   rD   rM   rN   s   @r3   rÎ   rÎ     sD   ø„ ñð
>˜{ð >¨tõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r4   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 )ÚIJepaAttentionr1   r7   Nc                 ó€   •— t         ‰| �  «        t        |«      | _        t	        |«      | _        t        «       | _        y rÐ   )r"   r#   r¬   Ú	attentionrÎ   ÚoutputÚsetÚpruned_headsr·   s     €r3   r#   zIJepaAttention.__init__1  s0   ø€ Ü‰ÑÔÜ+¨FÓ3ˆŒÜ% fÓ-ˆŒÜ›EˆÕr4   Ú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   ©rk   )Úlenr   rÙ   r¯   r±   rÜ   r   r›   rœ   r�   rÚ   rÑ   r²   Úunion)r0   rÝ   Úindexs      r3   Úprune_headszIJepaAttention.prune_heads7  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Ó:ˆÕr4   rÄ   r¼   r½   c                 óh   — | j                  |||«      }| j                  |d   |«      }|f|dd  z   }|S )Nr   r   )rÙ   rÚ   )r0   rÄ   r¼   r½   Úself_outputsÚattention_outputrÌ   s          r3   rD   zIJepaAttention.forwardI  sE   € ð —~‘~ m°YÐ@QÓRˆàŸ;™; |°A¡¸ÓFÐà#Ð%¨°Q°RÐ(8Ñ8ˆØˆr4   ry   )rF   rG   rH   r   r#   r   rz   rã   rJ   rK   r   rL   r	   r   rD   rM   rN   s   @r3   r×   r×   0  s’   ø„ ð"˜{ð "¨tõ "ð;  S¡ð ;¨dó ;ð* -1Ø"'ñ	à—|‘|ðð ˜EŸL™LÑ)ðð  ð	ð
 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷r4   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 )ÚIJepaIntermediater1   r7   Nc                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y rÐ   )r"   r#   r-   r…   r'   Úintermediate_sizerÑ   r(   Ú
hidden_actÚstrr   Úintermediate_act_fnr·   s     €r3   r#   zIJepaIntermediate.__init__X  s]   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3KÑ3KÓLˆŒ
Ü�f×'Ñ'¬Ô-Ü'-¨f×.?Ñ.?Ñ'@ˆDÕ$à'-×'8Ñ'8ˆDÕ$r4   rÄ   c                 óJ   — | j                  |«      }| j                  |«      }|S rÐ   )rÑ   rí   )r0   rÄ   s     r3   rD   zIJepaIntermediate.forward`  s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆàÐr4   ©	rF   rG   rH   r   r#   rJ   rK   rD   rM   rN   s   @r3   rè   rè   W  s1   ø„ ð9˜{ð 9¨tõ 9ð U§\¡\ð °e·l±l÷ r4   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 )ÚIJepaOutputr1   r7   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j                  «      | _	        y rÐ   )
r"   r#   r-   r…   rê   r'   rÑ   rY   rZ   r[   r·   s     €r3   r#   zIJepaOutput.__init__h  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×7Ñ7¸×9KÑ9KÓLˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r4   rÄ   rÒ   c                 óT   — | j                  |«      }| j                  |«      }||z   }|S rÐ   rÔ   rÕ   s      r3   rD   zIJepaOutput.forwardm  s.   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆà%¨Ñ4ˆàÐr4   rï   rN   s   @r3   rñ   rñ   g  s?   ø„ ð>˜{ð >¨tõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r4   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 )r   z?This corresponds to the Block class in the timm implementation.r1   r7   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-   r�   r'   Úlayer_norm_epsÚlayernorm_beforeÚlayernorm_afterr·   s     €r3   r#   zIJepaLayer.__init__y  s‡   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ'¨Ó/ˆŒÜ-¨fÓ5ˆÔÜ! &Ó)ˆŒÜ "§¡¨V×-?Ñ-?ÀV×EZÑEZÔ [ˆÔÜ!Ÿ|™|¨F×,>Ñ,>ÀF×DYÑDYÔZˆÕr4   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Ú   )r0   rÄ   r¼   r½   Úself_attention_outputsræ   rÌ   Úlayer_outputs           r3   rD   zIJepaLayer.forwardƒ  s–   € ð "&§¡Ø×!Ñ! -Ó0ØØ/ð "0ó "
Ðð
 2°!Ñ4ÐØ(¨¨Ð,ˆð )¨=Ñ8ˆð ×+Ñ+¨MÓ:ˆØ×(Ñ(¨Ó6ˆð —{‘{ <°Ó?ˆà�/ GÑ+ˆàˆr4   ry   )rF   rG   rH   rI   r   r#   rJ   rK   r   rL   r	   r   rD   rM   rN   s   @r3   r   r   v  s�   ø„ ÙIð[˜{ð [¨tõ [ð -1Ø"'ñ	à—|‘|ðð ˜EŸL™LÑ)ðð  ð	ð
 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷r4   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 )ÚIJepaEncoderr1   r7   Nc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w ry   )
r"   r#   r1   r-   Ú
ModuleListÚrangeÚnum_hidden_layersr   ÚlayerÚgradient_checkpointing)r0   r1   ru   r2   s      €r3   r#   zIJepaEncoder.__init__¡  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]ÄÀf×F^ÑF^Ó@_Ö#`¸1¤J¨vÕ$6Ò#`ÓaˆŒ
Ø&+ˆÕ#ùò $a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 )Nrš   r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrÐ   rš   )Ú.0Úvs     r3   ú	<genexpr>z'IJepaEncoder.forward.<locals>.<genexpr>Ë  s   è ø€ Òm˜qÐ_`Ñ_lœÑmùs   ‚Š)Úlast_hidden_staterÄ   Ú
attentions)Ú	enumerater  r  r£   Ú_gradient_checkpointing_funcÚ__call__Útupler   )r0   rÄ   r¼   r½   r	  r
  Úall_hidden_statesÚall_self_attentionsÚiÚlayer_moduleÚlayer_head_maskÚlayer_outputss               r3   rD   zIJepaEncoder.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ÜØ+Ø+Ø*ô
ð 	
r4   )NFFT)rF   rG   rH   r   r#   rJ   rK   r   rL   r	   r  r   rD   rM   rN   s   @r3   r  r     sz   ø„ ð,˜{ð ,¨tõ ,ð -1Ø"'Ø%*Ø ñ)
à—|‘|ð)
ð ˜EŸL™LÑ)ð)
ð  ð	)
ð
 #ð)
ð ð)
ð 
ˆu�oÐ%Ñ	&÷)
r4   r  c                   ó*   ‡ — e Zd Zdefˆ fd„Zd„ Zˆ xZS )ÚIJepaPoolerr1   c                 ó°   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                     | _	        y rÐ   )
r"   r#   r-   r…   r'   Úpooler_output_sizerÑ   r   Ú
pooler_actÚ
activationr·   s     €r3   r#   zIJepaPooler.__init__Ô  s>   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3LÑ3LÓMˆŒ
Ü  ×!2Ñ!2Ñ3ˆ�r4   c                 ó\   — |d d …df   }| j                  |«      }| j                  |«      }|S )Nr   )rÑ   r!  )r0   rÄ   Úfirst_token_tensorÚpooled_outputs       r3   rD   zIJepaPooler.forwardÙ  s6   € ð +ª1¨a¨4Ñ0ÐØŸ
™
Ð#5Ó6ˆØŸ™¨Ó6ˆØÐr4   )rF   rG   rH   r   r#   rD   rM   rN   s   @r3   r  r  Ó  s   ø„ ð4˜{õ 4ö
r4   r  aË  
    Args:
        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See [`IJepaImageProcessor.__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.
)r   é   i   aG  
    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 ([`IJepaConfig`]): 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.
z_The bare IJepa 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 )Ú
IJepaModelr1   Úadd_pooling_layerrQ   c                 ó  •— t         ‰| �  |«       || _        t        ||¬«      | _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _        |rt        |«      nd | _        | j                  «        y )N)rQ   rö   )r"   r#   r1   rP   rC   r  Úencoderr-   r�   r'   rû   Ú	layernormr  ÚpoolerÚ	post_init)r0   r1   r(  rQ   r2   s       €r3   r#   zIJepaModel.__init__  sk   ø€ Ü‰Ñ˜Ô ØˆŒÜ)¨&ÀÔPˆŒÜ# FÓ+ˆŒäŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÙ->”k &Ô)ÀDˆŒð 	�‰Õr4   r7   c                 ó.   — | j                   j                  S rÐ   )rC   rV   )r0   s    r3   Úget_input_embeddingszIJepaModel.get_input_embeddings  s   € Ø�‰×/Ñ/Ð/r4   Ú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ã   )r0   r0  r  rÝ   s       r3   Ú_prune_headszIJepaModel._prune_heads  sE   € ð
 +×0Ñ0Ó2ò 	C‰LˆE�5Ø�L‰L×Ñ˜uÑ%×/Ñ/×;Ñ;¸EÕBñ	Cr4   Úvision)Ú
checkpointÚoutput_typer“   ÚmodalityÚexpected_outputr5   ro   r¼   r½   r	  r6   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)ro   r6   )r¼   r½   r	  r
  r   r   )r  Úpooler_outputrÄ   r  )r1   r½   r	  Úuse_return_dictr=   Úget_head_maskr  rC   rV   r/   rˆ   r�   rŠ   r*  r+  r,  r   rÄ   r  )r0   r5   ro   r¼   r½   r	  r6   r
  Úexpected_dtypeÚembedding_outputÚencoder_outputsÚsequence_outputr$  Úhead_outputss                 r3   rD   zIJepaModel.forward%  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ô	
ð 	
r4   )FF©NNNNNNN)rF   rG   rH   r   rL   r#   r   r/  r   rz   r   r3  r   ÚIJEPA_INPUTS_DOCSTRINGr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCÚ_EXPECTED_OUTPUT_SHAPEr   rJ   rK   r{   r	   r   rD   rM   rN   s   @r3   r'  r'  	  s"  ø„ ñ

˜{ð 
¸tð 
Ð]aõ 
ð0Ð&:ó 0ðC¨4°°T¸#±Y°Ñ+?ð CÀDó Cñ +Ð+AÓBÙØ&Ø.Ø$ØØ.ôð 04Ø6:Ø,0Ø,0Ø/3Ø37Ø&*ñ;
à˜uŸ|™|Ñ,ð;
ð " %×"2Ñ"2Ñ3ð;
ð ˜EŸL™LÑ)ð	;
ð
 $ D™>ð;
ð ' t™nð;
ð #+¨4¡.ð;
ð ˜d‘^ð;
ð 
ˆuÐ0Ð0Ñ	1ò;
óó Cô;
r4   r'  zEgyptian cataÒ  
    IJepa Model transformer with an image classification head on top (a linear layer on top of the final hidden states)
    e.g. for ImageNet.

    <Tip>

        Note that it's possible to fine-tune IJepa 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 )ÚIJepaForImageClassificationr1   r7   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     €r3   r#   z$IJepaForImageClassification.__init__  ss   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ ¸%Ô@ˆŒ
ð OU×N_ÑN_ÐbcÒNcœ"Ÿ)™) F×$6Ñ$6¸×8IÑ8IÔJÔik×itÑitÓivˆŒð 	�‰Õr4   )r5  r6  r“   r8  r5   r¼   Úlabelsr½   r	  r6   r
  c                 ón  — |�|n| j                   j                  }| j                  ||||||¬«      }|d   }	| j                  |	j	                  d¬«      «      }
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	  r6   r
  r   r   rß   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationr]   )ÚlossÚlogitsrÄ   r  )r1   r;  r~   rL  rƒ   rŠ   ÚdeviceÚproblem_typerJ  r�   rJ   Úlongrz   r   Úsqueezer   rh   r
   r   rÄ   r  )r0   r5   r¼   rM  r½   r	  r6   r
  rÌ   r@  rS  rR  Úloss_fctrÚ   s                 r3   rD   z#IJepaForImageClassification.forward‹  sñ  € ð. &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—*‘*ØØØ/Ø!5Ø%=Ø#ð ó 
ˆð " !™*ˆà—‘ ×!5Ñ!5¸!Ð!5Ó!<Ó=ˆàˆØÑà—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ä$ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r4   rB  )rF   rG   rH   r   r#   r   rC  r   Ú_IMAGE_CLASS_CHECKPOINTr   rE  Ú_IMAGE_CLASS_EXPECTED_OUTPUTr   rJ   rK   rL   r	   r  rD   rM   rN   s   @r3   rH  rH  o  sï   ø„ ð 
˜{ð 
¨tõ 
ñ +Ð+AÓBÙØ*Ø)Ø$Ø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
óó CôA
r4   rH  )r}   r'  rH  )r‚   )@Úcollections.abcr)   Útypingr   r   r   r   r   r   r	   rJ   Útorch.nnr-   r
   r   r   Úactivationsr   Úmodeling_outputsr   r   r   Úmodeling_utilsr   r   Úpytorch_utilsr   r   Úutilsr   r   r   r   r   Úconfiguration_ijepar   Ú
get_loggerrF   rÂ   rD  rE  ÚModuler   rP   r}   rK   Úfloatrª   r¬   rÎ   r×   rè   rñ   r   r  r  rC  rF  ÚIJEPA_START_DOCSTRINGr'  rY  rZ  rH  Ú__all__rš   r4   r3   ú<module>ri     s  ðó ß D× DÑ Dã Ý ß AÑ Aå !ß bÑ bß Fß Q÷õ õ -ð 
ˆ×	Ñ	˜HÓ	%€ð 1Ð ð  €ô$˜2Ÿ9™9ô $ôNN�b—i‘iô Nôb"/˜?ô "/ðX ñ%Ø�I‰Ið%à�<‰<ð%ð 
�‰ð%ð �<‰<ð	%ð
 ˜UŸ\™\Ñ*ð%ð ð%ð ó%ô<;˜Ÿ™ô ;ô|�b—i‘iô ô$$�R—Y‘Yô $ôN˜Ÿ	™	ô ô �"—)‘)ô ô'�—‘ô 'ôT0
�2—9‘9ô 0
ôf�"—)‘)ô ðÐ ò2 (Ð ð	Ð ñ ØeØóô[
Ð%ó [
ó	ð[
ð| 5Ð Ø-Ð ñ ðð óôU
Ð"6ó U
óðU
òp P�r4   