§
    ‚Štj 4  ã                   óÂ  — d Z ddlmZ ddlZddlmZ ddlmZ ddl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 ddlmZmZmZmZmZ ddlmZ  G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z  ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z! ed¬¦  «         G d„ de¦  «        ¦   «         Z"g d ¢Z#dS )!z7PyTorch DeiT (Data-efficient Image Transformers) model.é    )Ú	dataclassN)Únné   )Úinitialization)ÚBaseModelOutputWithPoolingÚMaskedImageModelingOutput)ÚUnpack)ÚModelOutputÚTransformersKwargsÚauto_docstringÚ	torch_int)Úcan_return_tupleé   )ÚViTEmbeddingsÚViTForImageClassificationÚViTForMaskedImageModelingÚViTModelÚViTPreTrainedModelé   )Ú
DeiTConfigc            	       ó    ‡ — 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j        dz  dedej        fd„Zˆ xZS )ÚDeiTEmbeddingsaÇ  
    Construct the CLS token, distillation token, position and patch embeddings. Optionally, also the mask token.

    Differences from ViTEmbeddings:
    - Adds a distillation token (for distillation pre-training).
    - Position embeddings include +2 slots (CLS + distillation) instead of +1.
    - interpolate_pos_encoding handles 2 special tokens instead of 1.
    - forward concatenates distillation token and handles position encoding for both.
    FÚconfigÚuse_mask_tokenÚreturnNc                 ó˜  •— t          ¦   «                              ||¬¦  «         t          j        t	          j        dd|j        ¦  «        ¦  «        | _        | j        j	        }t          j        t	          j        d|dz   |j        ¦  «        ¦  «        | _
        t          j        t	          j        dd|j        ¦  «        ¦  «        | _        d S )N)r   r   r   )ÚsuperÚ__init__r   Ú	ParameterÚtorchÚzerosÚhidden_sizeÚ	cls_tokenÚpatch_embeddingsÚnum_patchesÚposition_embeddingsÚdistillation_token)Úselfr   r   r%   Ú	__class__s       €úc/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/deit/modular_deit.pyr   zDeiTEmbeddings.__init__2   sž   ø€ Ý‰Œ×Ò˜°ÐÑ?Ô?Ð?Ýœ¥e¤k°!°Q¸Ô8JÑ&KÔ&KÑLÔLˆŒØÔ+Ô7ˆå#%¤<µ´¸A¸{ÈQ¹ÐPVÔPbÑ0cÔ0cÑ#dÔ#dˆÔ Ý"$¤,­u¬{¸1¸aÀÔASÑ/TÔ/TÑ"UÔ"UˆÔÐÐó    Ú
embeddingsÚheightÚwidthc                 ó”  — |j         d         dz
  }| j        j         d         dz
  }t          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  ¦  «        }|                     d|||¦  «        }|                     dddd¦  «        }t          j
                             ||	|
fdd	¬
¦  «        }|                     dddd¦  «                             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)Úshaper&   r    ÚjitÚ
is_tracingÚ
patch_sizer   ÚreshapeÚpermuter   Ú
functionalÚinterpolateÚviewÚcat)r(   r,   r-   r.   r%   Únum_positionsÚclass_and_dist_pos_embedÚpatch_pos_embedr7   Ú
new_heightÚ	new_widthÚsqrt_num_positionss               r*   Úinterpolate_pos_encodingz'DeiTEmbeddings.interpolate_pos_encoding:   st  € ð !Ô& qÔ)¨AÑ-ˆØÔ0Ô6°qÔ9¸AÑ=ˆõ Œy×#Ò#Ñ%Ô%ð 	,¨+¸Ò*FÐ*FÈ6ÐUZÊ?È?ØÔ+Ð+à#'Ô#;¸A¸A¸A¸rÀ¸r¸EÔ#BÐ ØÔ2°1°1°1°a°b°b°5Ô9ˆàÔ˜rÔ"ˆà˜tœÑ.ˆ
Ø˜Tœ_Ñ,ˆ	å& }°cÑ'9Ñ:Ô:ÐØ)×1Ò1°!Ð5GÐI[Ð]`ÑaÔaˆØ)×1Ò1°!°Q¸¸1Ñ=Ô=ˆåœ-×3Ò3ØØ˜iÐ(ØØð	 4ñ 
ô 
ˆð *×1Ò1°!°Q¸¸1Ñ=Ô=×BÒBÀ1ÀbÈ#ÑNÔNˆåŒyÐ2°OÐDÈ!ÐLÑLÔLÐLr+   Úpixel_valuesÚbool_masked_posrH   c                 óâ  — |j         \  }}}}|                      |¦  «        }|                     ¦   «         \  }}	}|�R| j                             ||	d¦  «        }
|                     d¦  «                             |
¦  «        }|d|z
  z  |
|z  z   }| j                             |dd¦  «        }| j                             |dd¦  «        }t          j
        |||fd¬¦  «        }|r||                      |||¦  «        z   }n^|| j        d         k    s|| j        d         k    r2t          d|› d|› d| j        d         › d| j        d         › d	�	¦  «        ‚|| j        z   }|                      |¦  «        }|S )
Nr0   g      ð?r   r6   r   zInput image size (Ú*z) doesn't match model (z).)r8   r$   r3   Ú
mask_tokenÚexpandÚ	unsqueezeÚtype_asr#   r'   r    rA   rH   Ú
image_sizeÚ
ValueErrorr&   Údropout)r(   rI   rJ   rH   Ú_r-   r.   r,   Ú
batch_sizeÚ
seq_lengthÚmask_tokensÚmaskÚ
cls_tokensÚdistillation_tokenss                 r*   ÚforwardzDeiTEmbeddings.forwardb   s¾  € ð +Ô0Ñˆˆ1ˆf�eØ×*Ò*¨<Ñ8Ô8ˆ
à$.§O¢OÑ$5Ô$5Ñ!ˆ
�J àÐ&Øœ/×0Ò0°¸ZÈÑLÔLˆKà"×,Ò,¨RÑ0Ô0×8Ò8¸ÑEÔEˆDØ# s¨T¡zÑ2°[À4Ñ5GÑGˆJà”^×*Ò*¨:°r¸2Ñ>Ô>ˆ
Ø"Ô5×<Ò<¸ZÈÈRÑPÔPÐÝ”Y 
Ð,?ÀÐLÐRSÐTÑTÔTˆ
à#ð 	?Ø# d×&CÒ&CÀJÐPVÐX]Ñ&^Ô&^Ñ^ˆJˆJà˜œ¨Ô+Ò+Ð+¨u¸¼ÈÔ8JÒ/JÐ/JÝ ðE¨ð Eð E°%ð Eð EØœ¨Ô+ðEð EØ.2¬o¸aÔ.@ðEð Eð Eñô ð ð $ dÔ&>Ñ>ˆJà—\’\ *Ñ-Ô-ˆ
ØÐr+   )F)NF)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úboolr   r    ÚTensorÚintrH   Ú
BoolTensorr[   Ú__classcell__©r)   s   @r*   r   r   '   sþ   ø€ € € € € ðð ðVð V˜zð V¸4ð VÈDð Vð Vð Vð Vð Vð Vð&M°5´<ð &MÈð &MÐUXð &MÐ]bÔ]ið &Mð &Mð &Mð &MðV 48Ø).ð	 ð  à”lð ð Ô)¨DÑ0ð ð #'ð	 ð
 
Œð ð  ð  ð  ð  ð  ð  ð  r+   r   c                   óZ   ‡ — e Zd ZddgZdej        ej        z  ej        z  ddfˆ fd„Zˆ xZ	S )ÚDeiTPreTrainedModelr   Ú	DeiTLayerÚmoduler   Nc                 óR  •— t          ¦   «                              |¦  «         t          |t          ¦  «        rmt	          j        |j        ¦  «         t	          j        |j        ¦  «         t	          j        |j        ¦  «         |j	        �t	          j        |j	        ¦  «         dS dS dS )zInitialize the weightsN)
r   Ú_init_weightsÚ
isinstancer   ÚinitÚzeros_r#   r&   r'   rM   )r(   ri   r)   s     €r*   rk   z!DeiTPreTrainedModel._init_weightsˆ   sš   ø€ å‰Œ×Ò˜fÑ%Ô%Ð%Ý�f�nÑ-Ô-ð 	/ÝŒK˜Ô(Ñ)Ô)Ð)ÝŒK˜Ô2Ñ3Ô3Ð3ÝŒK˜Ô1Ñ2Ô2Ð2ØÔ Ð,Ý”˜FÔ-Ñ.Ô.Ð.Ð.Ð.ð	/ð 	/ð -Ð,r+   )
r\   r]   r^   Ú_no_split_modulesr   ÚLinearÚConv2dÚ	LayerNormrk   rd   re   s   @r*   rg   rg   …   sf   ø€ € € € € Ø)¨;Ð7Ðð/ B¤I°´	Ñ$9¸B¼LÑ$Hð /ÈTð /ð /ð /ð /ð /ð /ð /ð /ð /ð /r+   rg   c                   ó   — e Zd ZdS )Ú	DeiTModelN©r\   r]   r^   © r+   r*   rt   rt   “   ó   € € € € € Ø€Dr+   rt   c                   ó”   — e Zd Zee	 	 	 	 d
dej        dz  dej        dz  dedej        dz  de	e
         defd	„¦   «         ¦   «         ZdS )ÚDeiTForMaskedImageModelingNFrI   rJ   rH   Úattention_maskÚkwargsr   c                 ó6  —  | j         |f|||dœ|¤Ž}|j        }|dd…dd…f         }|j        \  }}	}
t          |	dz  ¦  «        x}}|                     ddd¦  «                             ||
||¦  «        }|                      |¦  «        }d}|�ñ| j        j        | j        j	        z  }|                     d||¦  «        }| 
                    | j        j	        d¦  «         
                    | j        j	        d¦  «                             d¦  «                             ¦   «         }t          j                             ||d¬	¦  «        }||z                       ¦   «         |                     ¦   «         d
z   z  | j        j        z  }t%          |||j        |j        ¬¦  «        S )a;  
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, num_patches)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).

        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]
        ```)rJ   rH   rz   Nr   r1   r   r   r0   Únone)Ú	reductiongñhãˆµøä>)ÚlossÚreconstructionÚhidden_statesÚ
attentions)ÚdeitÚlast_hidden_stater8   rb   r=   r<   Údecoderr   rQ   r;   Úrepeat_interleaverO   Ú
contiguousr   r>   Úl1_lossÚsumÚnum_channelsr   r�   r‚   )r(   rI   rJ   rH   rz   r{   ÚoutputsÚsequence_outputrU   Úsequence_lengthrŠ   r-   r.   Úreconstructed_pixel_valuesÚmasked_im_lossr3   rX   Úreconstruction_losss                     r*   r[   z"DeiTForMaskedImageModeling.forward˜   s¼  € ðL /8¨d¬iØð/
à+Ø%=Ø)ð	/
ð /
ð
 ð/
ð /
ˆð "Ô3ˆð *¨!¨!¨!¨Q¨R¨R¨%Ô0ˆØ4CÔ4IÑ1ˆ
�O \Ý˜_¨cÑ1Ñ2Ô2Ð2ˆ�Ø)×1Ò1°!°Q¸Ñ:Ô:×BÒBÀ:È|Ð]cÐejÑkÔkˆð &*§\¢\°/Ñ%BÔ%BÐ"àˆØÐ&Ø”;Ô)¨T¬[Ô-CÑCˆDØ-×5Ò5°b¸$ÀÑEÔEˆOà×1Ò1°$´+Ô2HÈ!ÑLÔLß"Ò" 4¤;Ô#9¸1Ñ=Ô=ß’˜1‘”ß’‘”ð	 õ #%¤-×"7Ò"7¸ÐF`ÐlrÐ"7Ñ"sÔ"sÐØ1°DÑ8×=Ò=Ñ?Ô?À4Ç8Â8Á:Ä:ÐPTÑCTÑUÐX\ÔXcÔXpÑpˆNå(ØØ5Ø!Ô/ØÔ)ð	
ñ 
ô 
ð 	
r+   )NNFN)r\   r]   r^   r   r   r    ra   rc   r`   r	   r   r   r[   rv   r+   r*   ry   ry   —   s³   € € € € € ØØð -1Ø37Ø).Ø.2ðJ
ð J
à”l TÑ)ðJ
ð Ô)¨DÑ0ðJ
ð #'ð	J
ð
 œ tÑ+ðJ
ð Ð+Ô,ðJ
ð 
#ðJ
ð J
ð J
ñ „^ñ ÔðJ
ð J
ð J
r+   ry   c                   ó   — e Zd ZdS )ÚDeiTForImageClassificationNru   rv   r+   r*   r’   r’   ç   rw   r+   r’   zC
    Output type of [`DeiTForImageClassificationWithTeacher`].
    )Úcustom_introc                   óÂ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
eej                 dz  ed<   dZeej                 dz  ed<   dS )Ú+DeiTForImageClassificationWithTeacherOutputaj  
    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).
    NÚlogitsÚ
cls_logitsÚdistillation_logitsr�   r‚   )r\   r]   r^   r_   r–   r    ÚFloatTensorÚ__annotations__r—   r˜   r�   Útupler‚   rv   r+   r*   r•   r•   ë   s¡   € € € € € € ð	ð 	ð (,€FˆEÔ Ñ$Ð+Ð+Ñ+Ø+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø48Ð˜Ô*¨TÑ1Ð8Ð8Ñ8Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø26€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6Ð6Ð6r+   r•   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	 	 	 ddej        dz  de	dej        dz  d	e
e         def
d
„¦   «         ¦   «         Zˆ xZS )Ú%DeiTForImageClassificationWithTeacherr   r   Nc                 ó¾  •— t          ¦   «                              |¦  «         |j        | _        t          |d¬¦  «        | _        |j        dk    rt          j        |j        |j        ¦  «        nt          j        ¦   «         | _	        |j        dk    rt          j        |j        |j        ¦  «        nt          j        ¦   «         | _
        |                      ¦   «          d S )NF)Úadd_pooling_layerr   )r   r   Ú
num_labelsrt   rƒ   r   rp   r"   ÚIdentityÚcls_classifierÚdistillation_classifierÚ	post_init)r(   r   r)   s     €r*   r   z.DeiTForImageClassificationWithTeacher.__init__  sÍ   ø€ Ý‰Œ×Ò˜Ñ Ô Ð à Ô+ˆŒÝ˜f¸Ð>Ñ>Ô>ˆŒ	ð AGÔ@QÐTUÒ@UÐ@U�BŒI�fÔ(¨&Ô*;Ñ<Ô<Ð<Õ[]Ô[fÑ[hÔ[hð 	Ôð AGÔ@QÐTUÒ@UÐ@U�BŒI�fÔ(¨&Ô*;Ñ<Ô<Ð<Õ[]Ô[fÑ[hÔ[hð 	Ô$ð
 	�ŠÑÔÐÐÐr+   FrI   rH   rz   r{   c                 ó  —  | j         |f||dœ|¤Ž}|j        }|                      |d d …dd d …f         ¦  «        }|                      |d d …dd d …f         ¦  «        }||z   dz  }	t	          |	|||j        |j        ¬¦  «        S )N)rH   rz   r   r   r   )r–   r—   r˜   r�   r‚   )rƒ   r„   r¢   r£   r•   r�   r‚   )
r(   rI   rH   rz   r{   r‹   rŒ   r—   r˜   r–   s
             r*   r[   z-DeiTForImageClassificationWithTeacher.forward!  sÌ   € ð /8¨d¬iØð/
à%=Ø)ð/
ð /
ð ð	/
ð /
ˆð "Ô3ˆà×(Ò(¨¸¸¸¸A¸q¸q¸q¸Ô)AÑBÔBˆ
Ø"×:Ò:¸?È1È1È1ÈaÐQRÐQRÐQRÈ7Ô;SÑTÔTÐð Ð2Ñ2°aÑ7ˆå:ØØ!Ø 3Ø!Ô/ØÔ)ð
ñ 
ô 
ð 	
r+   )NFN)r\   r]   r^   r   r   r   r   r    ra   r`   r	   r   r•   r[   rd   re   s   @r*   r�   r�     sÊ   ø€ € € € € ð˜zð ¨dð ð ð ð ð ð ð" Øð -1Ø).Ø.2ð	
ð 
à”l TÑ)ð
ð #'ð
ð œ tÑ+ð	
ð
 Ð+Ô,ð
ð 
5ð
ð 
ð 
ñ „^ñ Ôð
ð 
ð 
ð 
ð 
r+   r�   )r’   r�   ry   rt   rg   )$r_   Údataclassesr   r    r   Ú r   rm   Úmodeling_outputsr   r   Úprocessing_utilsr	   Úutilsr
   r   r   r   Úutils.genericr   Úvit.modeling_vitr   r   r   r   r   Úconfiguration_deitr   r   rg   rt   ry   r’   r•   r�   Ú__all__rv   r+   r*   ú<module>r¯      sÀ  ðð >Ð =à !Ð !Ð !Ð !Ð !Ð !à €€€Ø Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &ðð ð ð ð ð ð ð ð 'Ð &Ð &Ð &Ð &Ð &Ø OÐ OÐ OÐ OÐ OÐ OÐ OÐ OÐ OÐ OÐ OÐ OØ -Ð -Ð -Ð -Ð -Ð -ðð ð ð ð ð ð ð ð ð ð ð ð ð ð +Ð *Ð *Ð *Ð *Ð *ð[ð [ð [ð [ð [�]ñ [ô [ð [ð|/ð /ð /ð /ð /Ð,ñ /ô /ð /ð	ð 	ð 	ð 	ð 	�ñ 	ô 	ð 	ðM
ð M
ð M
ð M
ð M
Ð!:ñ M
ô M
ð M
ð`	ð 	ð 	ð 	ð 	Ð!:ñ 	ô 	ð 	ð €ððñ ô ð
 ð7ð 7ð 7ð 7ð 7°+ñ 7ô 7ñ „ñô ð7ð& €ðð
ñ 
ô 
ð0
ð 0
ð 0
ð 0
ð 0
Ð,?ñ 0
ô 0
ñ
ô 
ð0
ðfð ð €€€r+   