Ë
    T^(hX(  ã                   ó0  — d dl 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 ddlmZ ddlmZmZ d	d
lmZmZmZ dZ G d„ de«      Z G d„ de«      Zg d¢ZdZ ede«       G d„ dee«      «       ZdZdZ ede«       G d„ dee«      «       Zg d¢Zy)é    )ÚOptionalÚUnionN)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELoss)ÚIJepaConfigé   )ÚImageClassifierOutput)ÚPreTrainedModel)Úadd_start_docstringsÚ	torch_inté   )ÚViTEmbeddingsÚViTForImageClassificationÚViTModelzfacebook/ijepa_vith14_1kc            	       óÎ   ‡ — e 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 )ÚIJepaEmbeddingsÚconfigÚuse_mask_tokenÚreturnNc                 óÈ   •— t         ‰| �  ||«       | `| j                  j                  }t        j                  t        j                  d||j                  «      «      | _
        y )Né   )ÚsuperÚ__init__Ú	cls_tokenÚpatch_embeddingsÚnum_patchesÚnnÚ	ParameterÚtorchÚrandnÚhidden_sizeÚposition_embeddings)Úselfr   r   r   Ú	__class__s       €úe/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/ijepa/modular_ijepa.pyr   zIJepaEmbeddings.__init__   sL   ø€ Ü‰Ñ˜ Ô0àˆNØ×+Ñ+×7Ñ7ˆÜ#%§<¡<´·±¸A¸{ÈF×L^ÑL^Ó0_Ó#`ˆÕ ó    Ú
embeddingsÚheightÚwidthc                 ó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)Úshaper#   r    ÚjitÚ
is_tracingÚ
patch_sizer   ÚreshapeÚpermuter   Ú
functionalÚinterpolateÚview)r$   r(   r)   r*   r   Únum_positionsÚpatch_pos_embedÚdimÚ
new_heightÚ	new_widthÚsqrt_num_positionss              r&   Úinterpolate_pos_encodingz(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ˆàÐr'   Úpixel_valuesÚbool_masked_posr@   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)r@   r   r,   ç      ð?)	r1   r   Ú
mask_tokenÚexpandÚ	unsqueezeÚtype_asr@   r#   Údropout)r$   rA   rB   r@   Ú
batch_sizeÚ_r)   r*   r(   Ú
seq_lengthÚmask_tokensÚmasks               r&   ÚforwardzIJepaEmbeddings.forwardH   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à—\‘\ *Ó-ˆ
àÐr'   )F)NF)Ú__name__Ú
__module__Ú__qualname__r   Úboolr   r    ÚTensorÚintr@   r   Ú
BoolTensorrO   Ú__classcell__©r%   s   @r&   r   r      s•   ø„ ña˜{ð a¸Dð aÈTõ að%°5·<±<ð %Èð %ÐUXð %Ð]b×]iÑ]ió %ðT 7;Ø).ñ	à—l‘lðð " %×"2Ñ"2Ñ3ðð #'ð	ð
 
�‰÷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	)ÚIJepaPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚijeparA   Tr   Ú
IJepaLayerÚmoduler   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 weightsg        )ÚmeanÚstdNrD   )Ú
isinstancer   ÚLinearÚConv2dÚinitÚtrunc_normal_ÚweightÚdataÚtor    Úfloat32r   Úinitializer_rangeÚdtypeÚbiasÚzero_Ú	LayerNormÚfill_r   r#   rE   )r$   r]   s     r&   Ú_init_weightsz"IJepaPreTrainedModel._init_weightsq   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ð	 ×&Ñ&Ô+ð
 × Ñ Ð,Ø×!Ñ!×&Ñ&×,Ñ,Õ.ð -ð 1r'   )rP   rQ   rR   Ú__doc__r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_supports_sdpaÚ_supports_flash_attn_2r   r   rb   rc   rn   rp   © r'   r&   rZ   rZ   c   sa   „ ñð
 €LØÐØ$€OØ&*Ð#Ø*¨LÐ9ÐØ€NØ!Ðð/ E¨"¯)©)°R·Y±YÀÇÁÐ*LÑ$Mð /ÐRVô /r'   rZ   )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                   ó.   ‡ — e Zd Zddededefˆ fd„Zˆ xZS )Ú
IJepaModelr   Úadd_pooling_layerr   c                 óV   •— t         ‰| �  |«       || _        t        ||¬«      | _        y )N)r   )r   r   r   r   r(   )r$   r   r}   r   r%   s       €r&   r   zIJepaModel.__init__›   s%   ø€ Ü‰Ñ˜Ô ØˆŒÜ)¨&ÀÔPˆ�r'   )FF)rP   rQ   rR   r   rS   r   rW   rX   s   @r&   r|   r|   –   s(   ø„ ñ
Q˜{ð Q¸tð QÐ]a÷ Qñ Qr'   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fˆ fd„Z	 	 	 	 	 	 	 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 )ÚIJepaForImageClassificationr   c                 óh   •— t         ‰| �  |«       t        |d¬«      | _        | j	                  «        y )NF)r}   )r   r   r|   r[   Ú	post_init)r$   r   r%   s     €r&   r   z$IJepaForImageClassification.__init__µ   s(   ø€ Ü‰Ñ˜Ô Ü ¸%Ô@ˆŒ
Ø�‰Õr'   rA   Ú	head_maskÚlabelsÚoutput_attentionsÚoutput_hidden_statesr@   Úreturn_dictr   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†   r@   r‡   r   r   )r<   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationr,   )ÚlossÚlogitsÚhidden_statesÚ
attentions)r   Úuse_return_dictr[   Ú
classifierr_   rh   ÚdeviceÚproblem_typeÚ
num_labelsrk   r    ÚlongrU   r   Úsqueezer   r9   r   r
   rŽ   r�   )r$   rA   rƒ   r„   r…   r†   r@   r‡   ÚoutputsÚsequence_outputr�   rŒ   Úloss_fctÚoutputs                 r&   rO   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ä$ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r'   )NNNNNNN)rP   rQ   rR   r   r   r   r    rT   rS   r   Útupler
   rO   rW   rX   s   @r&   r€   r€   ¥   s¸   ø„ ð ˜{õ ð 04Ø,0Ø)-Ø,0Ø/3Ø37Ø&*ñA
à˜uŸ|™|Ñ,ðA
ð ˜EŸL™LÑ)ðA
ð ˜Ÿ™Ñ&ð	A
ð
 $ D™>ðA
ð ' t™nðA
ð #+¨4¡.ðA
ð ˜d‘^ðA
ð 
ˆuÐ+Ð+Ñ	,÷A
r'   r€   )rZ   r|   r€   ) Útypingr   r   r    Útorch.nnr   r   r   r   Ú-transformers.models.ijepa.configuration_ijepar   Úmodeling_outputsr
   Úmodeling_utilsr   Úutilsr   r   Úvit.modeling_vitr   r   r   Ú_CHECKPOINT_FOR_DOCr   rZ   Ú_EXPECTED_OUTPUT_SHAPEÚIJEPA_START_DOCSTRINGr|   Ú_IMAGE_CLASS_CHECKPOINTÚ_IMAGE_CLASS_EXPECTED_OUTPUTr€   Ú__all__ry   r'   r&   ú<module>r©      sÓ   ðß "ã Ý ß AÑ Aå Eå 5Ý -÷÷ñ ð 1Ð ôG�mô GôT"/˜?ô "/òJ (Ð ð	Ð ñ ØeØóôQÐ% xó Qó	ðQð 5Ð Ø-Ð ñ ðð óôG
Ð"6Ð8Qó G
óðG
òT�r'   