Ë
    S^(hçÝ  ã                   ó–  — d Z ddlZddlZddlmZ ddlmZ ddl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 ddl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!jJ                  e&«      Z'dZ(dZ)e G d„ de«      «       Z*e G d„ de«      «       Z+e G d„ de«      «       Z,dFd„Z-dGd„Z.dHd„Z/ G d„ dej`                  «      Z1 G d„ dej`                  «      Z2 G d„ d ej`                  «      Z3 G d!„ d"ej`                  «      Z4 G d#„ d$ej`                  «      Z5 G d%„ d&ej`                  «      Z6 G d'„ d(ej`                  «      Z7 G d)„ d*ej`                  «      Z8 G d+„ d,ej`                  «      Z9 G d-„ d.ej`                  «      Z: G d/„ d0ej`                  «      Z; G d1„ d2e«      Z<d3Z=d4Z> ed5e=«       G d6„ d7e<«      «       Z? G d8„ d9ej`                  «      Z@ ed:e=«       G d;„ d<e<«      «       ZA G d=„ d>ej`                  «      ZB G d?„ d@ej`                  «      ZC G dA„ dBej`                  «      ZD edCe=«       G dD„ dEe<«      «       ZEy)IzPyTorch TVLT model.é    N)Údeepcopy)Ú	dataclass)ÚOptionalÚTupleÚUnion)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )ÚACT2FN)ÚBaseModelOutputÚSequenceClassifierOutput)ÚPreTrainedModel)Ú find_pruneable_heads_and_indicesÚprune_linear_layer)ÚModelOutputÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstringsé   )Ú
TvltConfigr   zZinengTang/tvlt-basec                   óŽ  — 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j                     ed<   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                  d
f      ed<   dZeeej                  d
f      ed<   y)ÚTvltModelOutputaÃ  
    Class for TvltModel's outputs, with potential hidden states and attentions.

    Args:
        last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
            Sequence of hidden-states at the output of the last layer of the model.
        last_pixel_hidden_state (`torch.FloatTensor` of shape `(batch_size, pixel_sequence_length, hidden_size)`):
            Pixel sequence of hidden-states at the output of the last layer of the model.
        last_audio_hidden_state (`torch.FloatTensor` of shape `(batch_size, audio_sequence_length, hidden_size)`):
            Audio sequence of hidden-states at the output of the last layer of the model.
        pixel_label_masks (`torch.FloatTensor` of shape `(batch_size, pixel_patch_length)`):
            Tensor indicating which pixel patches are masked (1) and which are not (0).
        audio_label_masks (`torch.FloatTensor` of shape `(batch_size, audio_patch_length)`):
            Tensor indicating which audio patches are masked (1) and which are not (0).
        pixel_ids_restore (`torch.LongTensor` of shape `(batch_size, pixel_patch_length)`):
            Tensor containing the ids permutation of pixel masking.
        audio_ids_restore (`torch.LongTensor` of shape `(batch_size, audio_patch_length)`):
            Tensor containing the ids permutation of audio masking.
        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 and 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.
    NÚlast_hidden_stateÚlast_pixel_hidden_stateÚlast_audio_hidden_stateÚpixel_label_masksÚaudio_label_masksÚpixel_ids_restoreÚaudio_ids_restore.Úhidden_statesÚ
attentions)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   ÚtorchÚFloatTensorÚ__annotations__r   r   r   Ú
LongTensorr    r!   r"   r#   r   r$   © ó    úo/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/deprecated/tvlt/modeling_tvlt.pyr   r   0   sá   … ñð8 6:Ð�x × 1Ñ 1Ñ2Ó9Ø;?Ð˜X e×&7Ñ&7Ñ8Ó?Ø;?Ð˜X e×&7Ñ&7Ñ8Ó?Ø48Ð�x × 0Ñ 0Ñ1Ó8Ø48Ð�x × 0Ñ 0Ñ1Ó8Ø48Ð�x × 0Ñ 0Ñ1Ó8Ø48Ð�x × 0Ñ 0Ñ1Ó8Ø=A€M�8˜E %×"3Ñ"3°SÐ"8Ñ9Ñ:ÓAØ:>€J�˜˜u×0Ñ0°#Ð5Ñ6Ñ7Ô>r.   r   c                   óž   — e Zd ZU dZdZeej                     ed<   dZ	ee
ej                  df      ed<   dZee
ej                  df      ed<   y)ÚTvltDecoderOutputaM  
    Class for TvltDecoder's outputs, with potential hidden states and attentions.

    Args:
        logits (`torch.FloatTensor` of shape `(batch_size, patch_size ** 2 * num_channels)`):
            Pixel reconstruction logits.
        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 and 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.
    NÚlogits.r#   r$   )r%   r&   r'   r(   r2   r   r)   r*   r+   r#   r   r$   r-   r.   r/   r1   r1   Y   s\   … ñð  +/€FˆH�U×&Ñ&Ñ'Ó.Ø=A€M�8˜E %×"3Ñ"3°SÐ"8Ñ9Ñ:ÓAØ:>€J�˜˜u×0Ñ0°#Ð5Ñ6Ñ7Ô>r.   r1   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j                     ed<   dZeeej                  df      ed<   dZeeej                  df      ed	<   y)
ÚTvltForPreTrainingOutputa
  
    Class for TvltForPreTraining's outputs, with potential hidden states and attentions.

    Args:
        loss (`torch.FloatTensor` of shape `(1,)`):
            Pixel reconstruction loss.
        matching_logits (`torch.FloatTensor` of shape `(batch_size, 1)`):
            Matching objective logits.
        pixel_logits (`torch.FloatTensor` of shape
            `(batch_size, pixel_patch_length, image_patch_size ** 3 * pixel_num_channels)`): Pixel reconstruction
            logits.
        audio_logits (`torch.FloatTensor` of shape
            `(batch_size, audio_patch_length, image_patch_size[0] * image_patch_size[1])`): Audio reconstruction
            logits.
        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 and 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.
    NÚlossÚmatching_logitsÚpixel_logitsÚaudio_logits.r#   r$   )r%   r&   r'   r(   r5   r   r)   r*   r+   r6   r7   r8   r#   r   r$   r-   r.   r/   r4   r4   p   s›   … ñð0 )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø37€O�X˜e×/Ñ/Ñ0Ó7Ø04€L�(˜5×,Ñ,Ñ-Ó4Ø04€L�(˜5×,Ñ,Ñ-Ó4Ø=A€M�8˜E %×"3Ñ"3°SÐ"8Ñ9Ñ:ÓAØ:>€J�˜˜u×0Ñ0°#Ð5Ñ6Ñ7Ô>r.   r4   c                 ó–   — | j                   dd \  }}t        j                  ||f| j                  ¬«      }t	        |d|z
  z  «      }||fS )ú!Generate noise for audio masking.Né   ©Údevicer   )Úshaper)   Úrandr=   Úint)Úpixel_valuesÚ
pixel_maskÚ
mask_ratioÚ
batch_sizeÚseq_lenÚnoiseÚlen_keeps          r/   Úgenerate_pixel_mask_noiserH   ’   sS   € ð '×,Ñ,¨R¨aÐ0Ñ€J�Ü�J‰J˜
 GÐ,°\×5HÑ5HÔI€EÜ�7˜a *™nÑ-Ó.€HØ�(ˆ?Ðr.   c                 óX  — | j                   dd \  }}|dk(  rX||z  }t        j                  ||| j                  ¬«      j	                  d«      j                  dd|«      j                  ||«      }n'|dk(  r"t        j                  ||| j                  ¬«      }t        |d|z
  z  «      }	|	fS )r:   Nr;   zframe-levelr<   éÿÿÿÿr   úpatch-level)r>   r)   r?   r=   Ú	unsqueezeÚrepeatÚviewr@   )
Úaudio_valuesÚ
audio_maskrC   Ú	mask_typeÚfreq_lenrD   rE   Únum_time_patchesrF   rG   s
             r/   Úgenerate_audio_mask_noiserT   ›   s¬   € ð '×,Ñ,¨R¨aÐ0Ñ€J�Ø�MÒ!Ø" hÑ.Ðä�J‰J�zÐ#3¸L×<OÑ<OÔPß‰Y�r‹]ß‰V�A�q˜(Ó#ß‰T�*˜gÓ&ñ	 	ð 
�mÒ	#Ü—
‘
˜: w°|×7JÑ7JÔKˆÜ�7˜a *™nÑ-Ó.€HØ�(ˆ?Ðr.   c           	      óÚ  — | j                   \  }}}t        j                  |d¬«      }t        j                  |d¬«      }|dd…d|…f   }	t        j                  | d|	j	                  d«      j                  dd|«      ¬«      }
t        j                  ||g| j                  ¬«      }d|dd…d|…f<   t        j                  |d|¬«      }|�||z  }t        j                  |d|	¬«      }|
|||fS )z¸
    Perform random masking by per-sample shuffling on frame-level. Per-sample shuffling is done by argsort random
    noise. sequence: [batch_size, seq_len, hidden_dim], sequence
    r   ©ÚdimNrJ   ©rW   Úindexr<   r   )r>   r)   ÚargsortÚgatherrL   rM   Úonesr=   )ÚsequencerF   rG   Úattention_masksrD   rE   Ú
hidden_dimÚids_shuffleÚids_restoreÚids_keepÚsequence_maskedÚlabel_maskss               r/   Úrandom_maskingre   ­   sé   € ð '/§n¡nÑ#€J�˜ô —-‘- ¨1Ô-€KÜ—-‘- °Ô3€Kð š1˜i˜x˜i˜<Ñ(€HÜ—l‘l 8°¸(×:LÑ:LÈRÓ:P×:WÑ:WÐXYÐ[\Ð^hÓ:iÔj€Oô —*‘*˜j¨'Ð2¸8¿?¹?ÔK€KØ !€K’�9�H�9�Ñä—,‘,˜{°¸ÔE€KàÐ"Ø�Ñ&ˆÜŸ,™, ¸AÀXÔNˆà˜O¨[¸+ÐEÐEr.   c                   ó*   ‡ — e Zd ZdZˆ fd„Zdd„Zˆ xZS )ÚTvltPixelEmbeddingsú,Construct the patch and position embeddings.c                 ó  •— t         ‰| �  «        t        |«      | _        | j                  j                  | _        t        j                  t        j                  dd|j                  «      «      | _
        t        j                  t        j                  d|j                  |j                  «      «      | _        t        j                  t        j                  d| j                  |j                  «      «      | _        || _        y ©Nr   )ÚsuperÚ__init__ÚTvltPixelPatchEmbeddingsÚpatch_embeddingsÚnum_patches_per_imager   Ú	Parameterr)   ÚzerosÚhidden_sizeÚtype_embed_vÚ
num_framesÚtemporal_embedÚpos_embed_vÚconfig©Úselfrw   Ú	__class__s     €r/   rl   zTvltPixelEmbeddings.__init__Í   s¯   ø€ Ü‰ÑÔä 8¸Ó @ˆÔØ%)×%:Ñ%:×%PÑ%PˆÔ"äŸL™L¬¯©°Q¸¸6×;MÑ;MÓ)NÓOˆÔÜ Ÿl™l¬5¯;©;°q¸&×:KÑ:KÈV×M_ÑM_Ó+`ÓaˆÔÜŸ<™<¬¯©°A°t×7QÑ7QÐSY×SeÑSeÓ(fÓgˆÔàˆ�r.   c                 ó  — |j                   \  }}}}}| j                  |«      }|| j                  j                  d|d«      z  }|t	        j
                  | j                  d d …d |…f   | j                  d¬«      z  }|| j                  z  }||fS ©Nr   rV   )	r>   rn   rv   rM   r)   Úrepeat_interleaveru   ro   rs   )	ry   rA   r^   rD   rt   Únum_channelsÚheightÚwidthÚ
embeddingss	            r/   ÚforwardzTvltPixelEmbeddings.forwardÙ   s–   € à>J×>PÑ>PÑ;ˆ
�J ¨f°eà×*Ñ*¨<Ó8ˆ
Ø�d×&Ñ&×-Ñ-¨a°¸QÓ?Ñ?ˆ
Ø”e×-Ñ-¨d×.AÑ.AÂ!À[ÀjÀ[À.Ñ.QÐSW×SmÑSmÐstÔuÑuˆ
Ø�d×'Ñ'Ñ'ˆ
à˜?Ð*Ð*r.   ©N©r%   r&   r'   r(   rl   r‚   Ú__classcell__©rz   s   @r/   rg   rg   Ê   s   ø„ Ù6ô
÷	+r.   rg   c                   ó*   ‡ — e Zd ZdZˆ fd„Zdd„Zˆ xZS )ÚTvltAudioEmbeddingsrh   c                 ó¢  •— t         ‰| �  «        t        |«      | _        | j                  j                  | _        t        j                  t        j                  dd|j                  «      «      | _
        |j                  |j                  d   z  | _        t        j                  t        j                  d| j                  | j                  z  |j                  «      «      | _        t        j                  t        j                  d| j                  |j                  «      «      | _        |j                  |j                  d   z  | _        || _        y rj   )rk   rl   ÚTvltAudioPatchEmbeddingsrn   Únum_patchesr   rp   r)   rq   rr   Útype_embed_aÚfrequency_lengthÚaudio_patch_sizeÚnum_freq_patchesÚpos_embed_aÚ
freq_embedrw   rx   s     €r/   rl   zTvltAudioEmbeddings.__init__è   s÷   ø€ Ü‰ÑÔä 8¸Ó @ˆÔØ×0Ñ0×<Ñ<ˆÔäŸL™L¬¯©°Q¸¸6×;MÑ;MÓ)NÓOˆÔØ &× 7Ñ 7¸6×;RÑ;RÐSTÑ;UÑ UˆÔÜŸ<™<¬¯©°A°t×7GÑ7GÈ4×K`ÑK`Ñ7`Ðbh×btÑbtÓ(uÓvˆÔÜŸ,™,¤u§{¡{°1°d×6KÑ6KÈV×M_ÑM_Ó'`ÓaˆŒà &× 7Ñ 7¸6×;RÑ;RÐSTÑ;UÑ UˆÔØˆ�r.   c                 ó6  — | j                  |«      }|j                  d«      | j                  z  }|| j                  j	                  d|d«      z  }|t        j                  | j                  d d …d |…f   | j                  d¬«      z  }|| j                  z  }||fS r|   )	rn   Úsizer�   r‘   rM   r)   r}   r�   rŒ   )ry   rO   r^   r�   rS   s        r/   r‚   zTvltAudioEmbeddings.forwardö   s�   € à×*Ñ*¨<Ó8ˆ
à%Ÿ?™?¨1Ó-°×1FÑ1FÑFÐØ�d—o‘o×,Ñ,¨QÐ0@À!ÓDÑDˆ
Ø”e×-Ñ-¨d×.>Ñ.>ºqÐBSÐCSÐBSÐ?SÑ.TÐVZ×VkÑVkÐqrÔsÑsˆ
Ø�d×'Ñ'Ñ'ˆ
à˜?Ð*Ð*r.   rƒ   r„   r†   s   @r/   rˆ   rˆ   å   s   ø„ Ù6ô÷	+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 )rm   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)rk   rl   Ú
image_sizeÚimage_patch_sizeÚnum_image_channelsrr   Ú
isinstanceÚcollectionsÚabcÚIterableÚ
patch_sizer~   ro   r   ÚConv2dÚ
projection)ry   rw   r™   r    r~   rr   ro   rz   s          €r/   rl   z!TvltPixelPatchEmbeddings.__init__	  sß   ø€ Ü‰ÑÔØ!'×!2Ñ!2°F×4KÑ4K�Jˆ
Ø$*×$=Ñ$=¸v×?QÑ?Q�kˆä#-¨j¼+¿/¹/×:RÑ:RÔ#S‘ZÐZdÐfpÐYqˆ
Ü#-¨j¼+¿/¹/×:RÑ:RÔ#S‘ZÐZdÐfpÐYqˆ
Ø!+¨A¡°*¸Q±-Ñ!?ÀJÈqÁMÐU_Ð`aÑUbÑDbÑ cÐØ$ˆŒØ$ˆŒØ(ˆÔØ%:ˆÔ"Ø&ˆÔäŸ)™) L°+È:Ð^hÔiˆ�r.   rA   Úreturnc                 óì  — |j                   \  }}}}}|| j                  k7  rt        d«      ‚|| j                  d   k7  s|| j                  d   k7  r2t        d|› d|› d| j                  d   › d| j                  d   › d�	«      ‚|j	                  ||z  |||«      }| j                  |«      j                  d«      j                  dd«      }|j	                  ||| j                  z  | j                  «      }|S )	NúeMake sure that the channel dimension of the pixel values match with the one set in the configuration.r   r   zInput image size (Ú*ú) doesn't match model (ú).r;   )
r>   r~   Ú
ValueErrorr™   Úreshaper¢   ÚflattenÚ	transposero   rr   )ry   rA   rD   rt   r~   r   r€   r�   s           r/   r‚   z TvltPixelPatchEmbeddings.forward  s  € Ø>J×>PÑ>PÑ;ˆ
�J ¨f°eØ˜4×,Ñ,Ò,ÜØwóð ð �T—_‘_ QÑ'Ò'¨5°D·O±OÀAÑ4FÒ+FÜØ$ V H¨A¨e¨WÐ4KÈDÏOÉOÐ\]ÑL^ÐK_Ð_`Ðae×apÑapÐqrÑasÐ`tÐtvÐwóð ð $×+Ñ+¨J¸Ñ,CÀ\ÐSYÐ[`ÓaˆØ—_‘_ \Ó2×:Ñ:¸1Ó=×GÑGÈÈ1ÓMˆ
Ø×'Ñ'¨
°JÀ×A[ÑA[Ñ4[Ð]a×]mÑ]mÓnˆ
àÐr.   ©	r%   r&   r'   r(   rl   r)   ÚTensorr‚   r…   r†   s   @r/   rm   rm     s)   ø„ ñôjð  E§L¡Lð °U·\±\÷ r.   rm   c                   óZ   ‡ — e Zd ZdZˆ fd„Zdej                  dej                  fd„Zˆ xZS )rŠ   zì
    This class turns `audio_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
                  |j                  }}||f}t        |t        j                  j                  «      r|n||f}|d   |d   z  |d   |d   z  z  }|d   |d   z  |d   |d   z  f}	|| _        || _        || _        || _        |	| _        t!        j"                  ||||¬«      | _        y r–   )rk   rl   Úspectrogram_lengthr�   rŽ   Únum_audio_channelsrr   rœ   r�   rž   rŸ   Úspectrogram_sizer    r~   r‹   Úpatch_shaper   r¡   r¢   )ry   rw   r±   r�   r    r~   rr   r³   r‹   r´   rz   s             €r/   rl   z!TvltAudioPatchEmbeddings.__init__2  s  ø€ Ü‰ÑÔà×%Ñ%Ø×#Ñ#Ø×#Ñ#ð /9Ð,Ðð
 %+×$=Ñ$=¸v×?QÑ?Q�kˆà.Ð0@ÐAÐÜ#-¨j¼+¿/¹/×:RÑ:RÔ#S‘ZÐZdÐfpÐYqˆ
Ø'¨Ñ*¨j¸©mÑ;Ð@PÐQRÑ@SÐWaÐbcÑWdÑ@dÑeˆØ'¨Ñ*¨j¸©mÑ;Ð=MÈaÑ=PÐT^Ð_`ÑTaÑ=aÐbˆØ 0ˆÔØ$ˆŒØ(ˆÔØ&ˆÔØ&ˆÔäŸ)™) L°+È:Ð^hÔiˆ�r.   rO   r£   c                 óh  — |j                   \  }}}}|| j                  k7  rt        d«      ‚|| j                  d   kD  s|| j                  d   k7  r2t        d|› d|› d| j                  d   › d| j                  d   › d�	«      ‚| j	                  |«      j                  d«      j                  dd«      }|S )	Nr¥   r   r   zInput audio size (r¦   r§   r¨   r;   )r>   r~   r©   r³   r¢   r«   r¬   )ry   rO   rD   r~   r   r€   r�   s          r/   r‚   z TvltAudioPatchEmbeddings.forwardG  sÓ   € Ø2>×2DÑ2DÑ/ˆ
�L &¨%Ø˜4×,Ñ,Ò,ÜØwóð ð �D×)Ñ)¨!Ñ,Ò,°¸×9NÑ9NÈqÑ9QÒ0QÜØ$ V H¨A¨e¨Wð 5Ø×*Ñ*¨1Ñ-Ð.¨a°×0EÑ0EÀaÑ0HÐ/IÈðMóð ð —_‘_ \Ó2×:Ñ:¸1Ó=×GÑGÈÈ1ÓMˆ
àÐr.   r­   r†   s   @r/   rŠ   rŠ   +  s)   ø„ ñôjð* E§L¡Lð °U·\±\÷ r.   rŠ   c                   ó,   ‡ — e Zd Zˆ fd„Zd„ Zdd„Zˆ xZS )ÚTvltSelfAttentionc                 ó  •— t         ‰| �  «        |j                  |j                  z  dk7  r2t	        |d«      s&t        d|j                  › d|j                  › d�«      ‚|j                  | _        t        |j                  |j                  z  «      | _        | j                  | j                  z  | _        t        j                  |j                  | j                  |j                  ¬«      | _        t        j                  |j                  | j                  |j                  ¬«      | _        t        j                  |j                  | j                  |j                  ¬«      | _        t        j                  |j                   «      | _        y )Nr   Úembedding_sizezThe hidden size z4 is not a multiple of the number of attention heads ú.©Úbias)rk   rl   rr   Únum_attention_headsÚhasattrr©   r@   Úattention_head_sizeÚall_head_sizer   ÚLinearÚqkv_biasÚqueryÚkeyÚvalueÚDropoutÚattention_probs_dropout_probÚdropoutrx   s     €r/   rl   zTvltSelfAttention.__init__X  s.  ø€ Ü‰ÑÔØ×Ñ × :Ñ :Ñ:¸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ˆÔä—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Ô\ˆŒ
ä—z‘z &×"EÑ"EÓFˆ�r.   c                 ó    — |j                  «       d d | j                  | j                  fz   } |j                  |Ž }|j	                  dddd«      S )NrJ   r   r;   r   é   )r“   r½   r¿   rN   Úpermute)ry   ÚxÚnew_x_shapes      r/   Útranspose_for_scoresz&TvltSelfAttention.transpose_for_scoresj  sN   € Ø—f‘f“h˜s �m t×'?Ñ'?À×AYÑAYÐ&ZÑZˆØˆA�F‰F�KÐ ˆØ�y‰y˜˜A˜q !Ó$Ð$r.   c                 ó¶  — | j                  |«      }| j                  | j                  |«      «      }| j                  | j                  |«      «      }| j                  |«      }t	        j
                  ||j                  dd«      «      }	|	t        j                  | j                  «      z  }	|�|	|z   }	 t        j                  d¬«      |	«      }
| j                  |
«      }
|�|
|z  }
t	        j
                  |
|«      }|j                  dddd«      j                  «       }|j                  «       d d | j                   fz   } |j"                  |Ž }|r||
f}|S |f}|S )NrJ   éþÿÿÿrV   r   r;   r   rÊ   )rÃ   rÎ   rÄ   rÅ   r)   Úmatmulr¬   ÚmathÚsqrtr¿   r   ÚSoftmaxrÈ   rË   Ú
contiguousr“   rÀ   rN   )ry   r#   Úattention_maskÚ	head_maskÚoutput_attentionsÚmixed_query_layerÚ	key_layerÚvalue_layerÚquery_layerÚattention_scoresÚattention_probsÚcontext_layerÚnew_context_layer_shapeÚoutputss                 r/   r‚   zTvltSelfAttention.forwardo  sa  € Ø ŸJ™J }Ó5Ðà×-Ñ-¨d¯h©h°}Ó.EÓFˆ	Ø×/Ñ/°·
±
¸=Ó0IÓJˆØ×/Ñ/Ð0AÓBˆô !Ÿ<™<¨°Y×5HÑ5HÈÈRÓ5PÓQÐØ+¬d¯i©i¸×8PÑ8PÓ.QÑQÐØÐ%à/°.Ñ@Ðð -œ"Ÿ*™*¨Ô,Ð-=Ó>ˆð Ÿ,™, Ó7ˆð Ð Ø-°	Ñ9ˆOäŸ™ _°kÓBˆà%×-Ñ-¨a°°A°qÓ9×DÑDÓFˆØ"/×"4Ñ"4Ó"6°s¸Ð";¸t×?QÑ?QÐ>SÑ"SÐØ*˜×*Ñ*Ð,CÐDˆá6G�= /Ð2ˆàˆð O\ÐM]ˆàˆr.   ©NNF)r%   r&   r'   rl   rÎ   r‚   r…   r†   s   @r/   r·   r·   W  s   ø„ ôGò$%÷
!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 )	ÚTvltSelfOutputz¡
    The residual connection is defined in TvltLayer instead of here (as is the case with other models), due to the
    layernorm applied before each block.
    rw   r£   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  |j                  «      | _        y rƒ   )	rk   rl   r   rÁ   rr   ÚdenserÆ   Úhidden_dropout_probrÈ   rx   s     €r/   rl   zTvltSelfOutput.__init__™  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r.   r#   Úinput_tensorc                 óJ   — | j                  |«      }| j                  |«      }|S rƒ   ©ræ   rÈ   ©ry   r#   rè   s      r/   r‚   zTvltSelfOutput.forwardž  s$   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆàÐr.   )
r%   r&   r'   r(   r   rl   r)   r®   r‚   r…   r†   s   @r/   rä   rä   “  sD   ø„ ñð
>˜zð >¨dõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r.   rä   c                   ó,   ‡ — e Zd Zˆ fd„Zd„ Zdd„Zˆ xZS )ÚTvltAttentionc                 ó€   •— t         ‰| �  «        t        |«      | _        t	        |«      | _        t        «       | _        y rƒ   )rk   rl   r·   Ú	attentionrä   ÚoutputÚsetÚpruned_headsrx   s     €r/   rl   zTvltAttention.__init__¦  s0   ø€ Ü‰ÑÔÜ*¨6Ó2ˆŒÜ$ VÓ,ˆŒÜ›EˆÕr.   c                 ó>  — 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   rV   )Úlenr   rï   r½   r¿   rò   r   rÃ   rÄ   rÅ   rð   ræ   rÀ   Úunion)ry   ÚheadsrY   s      r/   Úprune_headszTvltAttention.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.   c                 ój   — | j                  ||||«      }| j                  |d   |«      }|f|dd  z   }|S )Nr   r   )rï   rð   )ry   r#   rÖ   r×   rØ   Úself_outputsÚattention_outputrá   s           r/   r‚   zTvltAttention.forward¾  sE   € Ø—~‘~ m°^ÀYÐPaÓbˆàŸ;™; |°A¡¸ÓFÐà#Ð%¨°Q°RÐ(8Ñ8ˆØˆr.   râ   )r%   r&   r'   rl   r÷   r‚   r…   r†   s   @r/   rí   rí   ¥  s   ø„ ô"ò;÷$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 )ÚTvltIntermediaterw   r£   Nc                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y rƒ   )rk   rl   r   rÁ   rr   Úintermediate_sizeræ   rœ   Ú
hidden_actÚstrr   Úintermediate_act_fnrx   s     €r/   rl   zTvltIntermediate.__init__È  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  ©ry   r#   s     r/   r‚   zTvltIntermediate.forwardÐ  s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆàÐr.   ©	r%   r&   r'   r   rl   r)   r®   r‚   r…   r†   s   @r/   rü   rü   Ç  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 )Ú
TvltOutputrw   r£   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j                  «      | _	        y rƒ   )
rk   rl   r   rÁ   rþ   rr   ræ   rÆ   rç   rÈ   rx   s     €r/   rl   zTvltOutput.__init__Ø  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/   r‚   zTvltOutput.forwardÝ  s.   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆà%¨Ñ4ˆàÐr.   r  r†   s   @r/   r  r  ×  s?   ø„ ð>˜zð >¨dõ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r.   r  c                   ó*   ‡ — e Zd ZdZˆ fd„Zdd„Zˆ xZS )Ú	TvltLayerz?This corresponds to the Block class in the timm implementation.c                 ór  •— t         ‰| �  «        |j                  | _        d| _        t	        |«      | _        t        |«      | _        t        |«      | _	        t        j                  |j                  |j                  ¬«      | _        t        j                  |j                  |j                  ¬«      | _        y ©Nr   ©Úeps)rk   rl   Úchunk_size_feed_forwardÚseq_len_dimrí   rï   rü   Úintermediater  rð   r   Ú	LayerNormrr   Úlayer_norm_epsÚlayernorm_beforeÚlayernorm_afterrx   s     €r/   rl   zTvltLayer.__init__é  s‡   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ& vÓ.ˆŒÜ,¨VÓ4ˆÔÜ  Ó(ˆŒÜ "§¡¨V×-?Ñ-?ÀV×EZÑEZÔ [ˆÔÜ!Ÿ|™|¨F×,>Ñ,>ÀF×DYÑDYÔZˆÕr.   c                 ó  — | j                  | j                  |«      |||¬«      }|d   }|dd  }||j                  |j                  «      z   }| j	                  |«      }| j                  |«      }| j                  ||«      }|f|z   }|S )N©rØ   r   r   )rï   r  Útor=   r  r  rð   )	ry   r#   rÖ   r×   rØ   Úself_attention_outputsrú   rá   Úlayer_outputs	            r/   r‚   zTvltLayer.forwardó  s©   € Ø!%§¡Ø×!Ñ! -Ó0ØØØ/ð	 "0ó "
Ðð 2°!Ñ4ÐØ(¨¨Ð,ˆð )¨=×+;Ñ+;Ð<L×<SÑ<SÓ+TÑTˆð ×+Ñ+¨MÓ:ˆØ×(Ñ(¨Ó6ˆð —{‘{ <°Ó?ˆà�/ GÑ+ˆàˆr.   râ   r„   r†   s   @r/   r
  r
  æ  s   ø„ ÙIô[÷r.   r
  c                   ó0   ‡ — e Zd Zˆ fd„Z	 	 	 	 	 dd„Zˆ xZS )ÚTvltEncoderc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w )NF)
rk   rl   rw   r   Ú
ModuleListÚrangeÚnum_hidden_layersr
  ÚlayerÚgradient_checkpointing)ry   rw   Ú_rz   s      €r/   rl   zTvltEncoder.__init__  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]¼uÀV×E]ÑE]Ó?^Ö#_¸!¤I¨fÕ$5Ò#_Ó`ˆŒ
Ø&+ˆÕ#ùò $`s   ½A#c                 óx  — |rdnd }|rdnd }t        | j                  «      D ]j  \  }	}
|r||fz   }|�||	   nd }| j                  r,| j                  r | j	                  |
j
                  ||||«      }n |
||||«      }|d   }|sŒb||d   fz   }Œl |r||fz   }|st        d„ |||fD «       «      S t        |||¬«      S )Nr-   r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrƒ   r-   ©Ú.0Úvs     r/   ú	<genexpr>z&TvltEncoder.forward.<locals>.<genexpr>9  s   è ø€ Òm˜qÐ_`Ñ_lœÑmùó   ‚Š)r   r#   r$   )Ú	enumerater!  r"  ÚtrainingÚ_gradient_checkpointing_funcÚ__call__Útupler   )ry   r#   rÖ   r×   rØ   Úoutput_hidden_statesÚreturn_dictÚall_hidden_statesÚall_self_attentionsÚiÚlayer_moduleÚlayer_head_maskÚlayer_outputss                r/   r‚   zTvltEncoder.forward  s  € ñ #7™B¸DÐÙ$5™b¸4Ðä(¨¯©Ó4ò 	P‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à.7Ð.C˜i¨šlÈˆOà×*Ò*¨t¯}ª}Ø $× AÑ AØ ×)Ñ)Ø!Ø"Ø#Ø%ó!‘ñ !-¨]¸NÈOÐ]nÓ o�à)¨!Ñ,ˆMâ Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð)	Pñ,  Ø 1°]Ð4DÑ DÐáÜÑm ]Ð4EÐGZÐ$[ÔmÓmÐmÜØ+Ø+Ø*ô
ð 	
r.   )NNFFT©r%   r&   r'   rl   r‚   r…   r†   s   @r/   r  r    s   ø„ ô,ð ØØØ"Ø÷+
r.   r  c                   ó&   — e Zd ZdZeZdZdZdZd„ Z	y)ÚTvltPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚtvltrA   Tc                 óú  — t        |t        j                  t        j                  f«      rm|j                  j
                  j                  d| j                  j                  ¬«       |j                  �%|j                  j
                  j                  «        yyt        |t        j                  «      rJ|j                  j
                  j                  «        |j                  j
                  j                  d«       yy)zInitialize the weightsç        )ÚmeanÚstdNg      ð?)rœ   r   rÁ   r¡   ÚweightÚdataÚnormal_rw   Úinitializer_ranger¼   Úzero_r  Úfill_)ry   Úmodules     r/   Ú_init_weightsz!TvltPreTrainedModel._init_weightsL  s¨   € ä�fœrŸy™y¬"¯)©)Ð4Ô5ð �M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡Ô-Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)ð .r.   N)
r%   r&   r'   r(   r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚsupports_gradient_checkpointingrG  r-   r.   r/   r:  r:  A  s$   „ ñð
 €LØÐØ$€OØ&*Ð#ó
*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 ([`TvltConfig`]): 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_frames, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`TvltProcessor`]. See [`TvltProcessor.__call__`] for
            details.

        audio_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Audio values. Audio values can be obtained using [`TvltProcessor`]. See [`TvltProcessor.__call__`] for
            details.

        pixel_mask (`torch.FloatTensor` of shape `(batch_size, num_pixel_patches)`):
            Pixel masks. Pixel masks can be obtained using [`TvltProcessor`]. See [`TvltProcessor.__call__`] for
            details.

        audio_mask (`torch.FloatTensor` of shape `(batch_size, num_audio_patches)`):
            Audio masks. Audio masks can be obtained using [`TvltProcessor`]. See [`TvltProcessor.__call__`] for
            details.

        pixel_values_mixed (`torch.FloatTensor` of shape `(batch_size, num_frames, num_channels, height, width)`):
            Pixel values that mix positive and negative samples in Tvlt vision-audio matching. Pixel values mixed can
            be obtained using [`TvltProcessor`]. See [`TvltProcessor.__call__`] for details.

        pixel_mask_mixed (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel masks of pixel_values_mixed. Pixel masks mixed can be obtained using [`TvltProcessor`]. See
            [`TvltProcessor.__call__`] for details.

        mask_pixel (`bool`, *optional*):
            Whether to mask pixel for MAE tasks. Only set to True in TvltForPreTraining.

        mask_audio (`bool`, *optional*):
            Whether to mask audio for MAE tasks. Only set to True in TvltForPreTraining.

        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.
z^The bare TVLT Model transformer outputting raw hidden-states without any specific head on top.c                   ó,  ‡ — e Zd Zˆ fd„Zd„ Zd„ Z ee«       ee	e
¬«      	 	 	 	 	 	 	 ddej                  dej                  deej                     deej                     d	ed
edee   dee   dee   deeej                     e	f   fd„«       «       Zˆ xZS )Ú	TvltModelc                 ó¬  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        t        |«      | _        t        j                  t        j                  dd|j                  «      «      | _        |j                  rd | _        n0t        j"                  |j                  |j$                  ¬«      | _        | j'                  «        y r  )rk   rl   rw   rg   Úpixel_embeddingsrˆ   Úaudio_embeddingsr  Úencoderr   rp   r)   rq   rr   Úcls_embeddingÚuse_mean_poolingÚ	layernormr  r  Ú	post_initrx   s     €r/   rl   zTvltModel.__init__–  s›   ø€ Ü‰Ñ˜Ô ØˆŒä 3°FÓ ;ˆÔÜ 3°FÓ ;ˆÔÜ" 6Ó*ˆŒäŸ\™\¬%¯+©+°a¸¸F×<NÑ<NÓ*OÓPˆÔà×"Ò"Ø!ˆD�NäŸ\™\¨&×*<Ñ*<À&×BWÑBWÔXˆDŒNð 	�‰Õr.   c                 óZ   — | j                   j                  | j                  j                  fS rƒ   )rO  rn   rP  )ry   s    r/   Úget_input_embeddingszTvltModel.get_input_embeddings¨  s%   € Ø×$Ñ$×5Ñ5°t×7LÑ7L×7]Ñ7]Ð]Ð]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)ÚitemsrQ  r!  rï   r÷   )ry   Úheads_to_pruner!  rö   s       r/   Ú_prune_headszTvltModel._prune_heads«  sE   € ð
 +×0Ñ0Ó2ò 	C‰LˆE�5Ø�L‰L×Ñ˜uÑ%×/Ñ/×;Ñ;¸EÕBñ	Cr.   ©Úoutput_typerH  rA   rO   rB   rP   Ú
mask_pixelÚ
mask_audiorØ   r0  r1  r£   c
                 óÜ  — |�|n| j                   j                  }|�|n| j                   j                  }|	�|	n| j                   j                  }	| j	                  ||«      \  }
}| j                  ||«      \  }}d}d}|r9t        |
|| j                   j                  ¬«      \  }}t        |
|||¬«      \  }
}}}d}d}|r| j                   j                  | j                   j                  d   z  }t        ||| j                   j                  | j                   j                  |¬«      \  }}t        ||||¬«      \  }}}}|j                  d«      }t        j                   | j"                  j%                  |dd«      |
|gd«      }|
j                  d«      }d}|�$|�"t        j                   |dd…dd…f   ||gd«      }|j                  «       }d}|�| j'                  ||«      }| j)                  |||||	¬«      }|d   }| j*                  �| j+                  |«      }|dd…dd|z   …f   }|dd…d|z   d…f   }|	s|||||||f|dd z   S t-        ||||||||j.                  |j0                  ¬«	      S )	a‡  
        Returns:

        Examples:

        ```python
        >>> from transformers import TvltProcessor, TvltModel
        >>> import numpy as np
        >>> import torch

        >>> num_frames = 8
        >>> images = list(np.random.randn(num_frames, 3, 224, 224))
        >>> audio = list(np.random.randn(10000))

        >>> processor = TvltProcessor.from_pretrained("ZinengTang/tvlt-base")
        >>> model = TvltModel.from_pretrained("ZinengTang/tvlt-base")

        >>> input_dict = processor(images, audio, sampling_rate=44100, return_tensors="pt")

        >>> outputs = model(**input_dict)
        >>> loss = outputs.loss
        ```N)rB   rC   )r^   r   )rP   rC   rQ   rR   r   )rÖ   rØ   r0  r1  )	r   r   r   r   r    r!   r"   r#   r$   )rw   rØ   r0  Úuse_return_dictrO  rP  rH   Úpixel_mask_ratiore   r�   rŽ   rT   Úaudio_mask_ratioÚaudio_mask_typer“   r)   ÚcatrR  rM   Úget_extended_attention_maskrQ  rT  r   r#   r$   )ry   rA   rO   rB   rP   r^  r_  rØ   r0  r1  Úpixel_embedding_outputÚaudio_embedding_outputr   r!   Úpixel_mask_noiseÚpixel_len_keepr    r"   r�   Úaudio_mask_noiseÚaudio_len_keeprD   Úembedding_outputÚmasked_pixel_lenrÖ   Úinput_shapeÚextended_attention_maskÚencoder_outputsÚsequence_outputÚpixel_sequence_outputÚaudio_sequence_outputs                                  r/   r‚   zTvltModel.forward³  s  € ðJ 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà-1×-BÑ-BÀ<ÐQ[Ó-\Ñ*Ð 
à-1×-BÑ-BÀ<ÐQ[Ó-\Ñ*Ð 
ð !ÐØ ÐÙÜ/HØ&°:È$Ï+É+×JfÑJfô0Ñ,Ð˜nô XfØ&Ø ØØ *ô	XÑTÐ" JÐ0AÐCTð !ÐØ ÐÙØ#Ÿ{™{×;Ñ;¸t¿{¹{×?[Ñ?[Ð\]Ñ?^Ñ^ÐÜ/HØ&Ø%ØŸ;™;×7Ñ7ØŸ+™+×5Ñ5Ø)ô0Ñ,Ð˜nô XfØ&Ø ØØ *ô	XÑTÐ" JÐ0AÐCTð "×&Ñ& qÓ)ˆ
Ü Ÿ9™9Ø×Ñ×&Ñ& z°1°aÓ8Ð:PÐRhÐiÐkló
Ðð 2×6Ñ6°qÓ9ÐàˆØÐ! jÐ&<Ü"ŸY™Y¨
²1°b°q°b°5Ñ(9¸:ÀzÐ'RÐTUÓVˆNà&×+Ñ+Ó-ˆØ"&ÐØÐ%Ø&*×&FÑ&FÀ~ÐWbÓ&cÐ#àŸ,™,ØØ2Ø/Ø!5Ø#ð 'ó 
ˆð *¨!Ñ,ˆØ�>‰>Ð%Ø"Ÿn™n¨_Ó=ˆOà /²°1°qÐ;KÑ7KÐ3KÐ0KÑ LÐØ /²°1Ð7GÑ3GÑ3IÐ0IÑ JÐÙàØ%Ø%Ø!Ø!Ø!Ø!ðð    Ð#ñ$ð $ô Ø-Ø$9Ø$9Ø/Ø/Ø/Ø/Ø)×7Ñ7Ø&×1Ñ1ô

ð 
	
r.   )NNFFNNN)r%   r&   r'   rl   rW  r[  r   ÚTVLT_INPUTS_DOCSTRINGr   r   Ú_CONFIG_FOR_DOCr)   r*   r   Úboolr   r   r‚   r…   r†   s   @r/   rM  rM  ‘  s  ø„ ô
ò$^òCñ +Ð+@ÓAÙ¨?ÈÔYð
 37Ø26Ø Ø Ø,0Ø/3Ø&*ñ@
à×'Ñ'ð@
ð ×'Ñ'ð@
ð ˜U×.Ñ.Ñ/ð	@
ð
 ˜U×.Ñ.Ñ/ð@
ð ð@
ð ð@
ð $ D™>ð@
ð ' t™nð@
ð ˜d‘^ð@
ð 
ˆu�U×&Ñ&Ñ'¨Ð8Ñ	9ò@
ó Zó Bô@
r.   rM  c                   ó,   ‡ — e Zd Zˆ fd„Z	 	 	 dd„Zˆ xZS )ÚTvltDecoderc                 óÎ  •— t         ‰| �  «        t        |«      }|j                  |_        |j
                  |_        |j                  |_        |j                  |_
        t        j                  t        |j
                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        t        j                   |j                  |j"                  ¬«      | _        d| _        || _        y c c}w )Nr  F)rk   rl   r   Údecoder_hidden_sizerr   Údecoder_num_hidden_layersr   Údecoder_num_attention_headsr½   Údecoder_intermediate_sizerþ   r   r  r  r
  Údecoder_layersr  r  rT  r"  rw   )ry   rw   Údecoder_configr#  rz   s       €r/   rl   zTvltDecoder.__init__9  s´   ø€ Ü‰ÑÔä! &Ó)ˆØ%+×%?Ñ%?ˆÔ"Ø+1×+KÑ+KˆÔ(Ø-3×-OÑ-OˆÔ*Ø+1×+KÑ+KˆÔ(Ü Ÿm™mÜ05°f×6VÑ6VÓ0WÖX¨1ŒY�~Õ&ÒXó
ˆÔô Ÿ™ f×&@Ñ&@Àf×F[ÑF[Ô\ˆŒà&+ˆÔ#Øˆ�ùò Ys   ÂC"c                 ó„  — |rdnd }|rdnd }t        | j                  «      D ]_  \  }}|r||fz   }| j                  r+| j                  r| j	                  |j
                  |d |«      }	n
 |||¬«      }	|	d   }|sŒW||	d   fz   }Œa |r||fz   }| j                  |«      }
|st        d„ |
||fD «       «      S t        |
||¬«      S )Nr-   r  r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrƒ   r-   r&  s     r/   r)  z&TvltDecoder.forward.<locals>.<genexpr>n  s   è ø€ Òf˜qÐXYÑXeœÑfùr*  )r2   r#   r$   )	r+  r  r"  r,  r-  r.  rT  r/  r1   )ry   r#   rØ   r0  r1  r2  r3  r4  r5  r7  r2   s              r/   r‚   zTvltDecoder.forwardJ  sú   € ñ #7™B¸DÐÙ$5™b¸4ÐÜ(¨×)<Ñ)<Ó=ò 	P‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à×*Ò*¨t¯}ª}Ø $× AÑ AØ ×)Ñ)Ø!ØØ%ó	!‘ñ !-¨]ÐN_Ô `�à)¨!Ñ,ˆMâ Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð#	Pñ&  Ø 1°]Ð4DÑ DÐð —‘ Ó.ˆáÜÑf VÐ->Ð@SÐ$TÔfÓfÐfÜ ¨Ð>OÐ\oÔpÐpr.   )FFTr8  r†   s   @r/   ry  ry  8  s   ø„ ôð(  Ø"Ø÷%qr.   ry  zTThe TVLT Model transformer with the decoder on top for self-supervised pre-training.c                   ó’  ‡ — e Zd Zˆ fd„Zd„ Zd„ Zd„ Zd„ Zd„ Z e	e
«       eee¬«      	 	 	 	 	 	 	 	 ddej                  d	ej                  d
eej                     deej                     deej"                     deej                     deej                     dee   dee   dee   deeej                     ef   fd„«       «       Zˆ xZS )ÚTvltForPreTrainingc                 ó–  •— t         ‰	| �  |«       || _        |j                  | _        |j                  | _        | j                  s| j                  st        d«      ‚t        |«      | _        | j                  rt        |«      | _	        | j                  �r$t        j                  |j                  |j                  d¬«      | _        t        j                  t!        j"                  dd|j                  «      «      | _        t        j                  t!        j"                  dd|j                  «      «      | _        t)        |«      | _        |j                  }|j,                  }| j                  j.                  j0                  }t        j                  t!        j"                  d||«      «      | _        t        j                  t!        j"                  d|j,                  |«      «      | _        t        j                  t!        j"                  dd|«      «      | _        | j                  j8                  j:                  }|j<                  |j>                  d   z  }t        j                  t!        j"                  d||z  |«      «      | _         t        j                  t!        j"                  d||«      «      | _!        t        j                  t!        j"                  dd|«      «      | _"        | j                  jF                  d   dz  | j                  jH                  z  }tK        ||«      | _&        | j                  j>                  d   | j                  j>                  d   z  | j                  jN                  z  }tK        ||«      | _(        || _        || _        || _)        |jF                  | _#        |j>                  | _        | jU                  «        y )Nz;Must set at least one of matching task and MAE task to trueTr»   r   r   r;   )+rk   rl   rw   Útask_matchingÚtask_maer©   rM  r;  ÚTvltMatchingHeadÚmatching_headr   rÁ   rr   r{  Úencoder_to_decoderrp   r)   rq   Úpixel_mask_tokenÚaudio_mask_tokenry  Údecoderrt   rO  ro   Údecoder_pixel_pos_embedÚdecoder_temporal_embedÚdecoder_pixel_type_embedrP  r‹   r�   rŽ   Údecoder_audio_pos_embedÚdecoder_freq_embedÚdecoder_audio_type_embedrš   r›   ÚTvltMAEHeadÚpixel_mae_headr²   Úaudio_mae_headr�   rU  )
ry   rw   r{  rt   ro   Únum_audio_patchesr�   Úpixel_mae_output_dimÚaudio_mae_output_dimrz   s
            €r/   rl   zTvltForPreTraining.__init__w  sß  ø€ Ü‰Ñ˜Ô ØˆŒà#×1Ñ1ˆÔØŸ™ˆŒØ×"Ò" d§m¢mÜÐZÓ[Ð[ä˜fÓ%ˆŒ	à×ÒÜ!1°&Ó!9ˆDÔà�=‹=Ü&(§i¡i°×0BÑ0BÀF×D^ÑD^ÐeiÔ&jˆDÔ#ä$&§L¡L´·±¸QÀÀ6×C]ÑC]Ó1^Ó$_ˆDÔ!Ü$&§L¡L´·±¸QÀÀ6×C]ÑC]Ó1^Ó$_ˆDÔ!ä& vÓ.ˆDŒLà"(×"<Ñ"<Ðà×*Ñ*ˆJØ$(§I¡I×$>Ñ$>×$TÑ$TÐ!Ü+-¯<©<¼¿¹ÀAÐG\Ð^qÓ8rÓ+sˆDÔ(Ü*,¯,©,´u·{±{À1Àf×FWÑFWÐYlÓ7mÓ*nˆDÔ'Ü,.¯L©L¼¿¹ÀQÈÐK^Ó9_Ó,`ˆDÔ)à $§	¡	× :Ñ :× FÑ FÐØ%×6Ñ6¸&×:QÑ:QÐRSÑ:TÑTÐÜ+-¯<©<Ü—‘˜AÐ0Ð4DÑDÐFYÓZó,ˆDÔ(ô ')§l¡l´5·;±;¸qÐBRÐTgÓ3hÓ&iˆDÔ#Ü,.¯L©L¼¿¹ÀQÈÐK^Ó9_Ó,`ˆDÔ)à#'§;¡;×#?Ñ#?ÀÑ#BÀaÑ#GÈ$Ï+É+×JhÑJhÑ#hÐ Ü"-¨fÐ6JÓ"KˆDÔà—‘×,Ñ,¨QÑ/°$·+±+×2NÑ2NÈqÑ2QÑQÐTX×T_ÑT_×TrÑTrÑrð !ô #.¨fÐ6JÓ"KˆDÔà(ˆDŒOØ)>ˆDÔ&Ø$4ˆDÔ!Ø$*×$;Ñ$;ˆDÔ!Ø$*×$;Ñ$;ˆDÔ!ð 	�‰Õr.   c           
      ó®  — |j                   \  }}}}}|j                   d   | j                  d   z  }|j                   d   | j                  d   z  }|j                  ||||| j                  d   || j                  d   f¬«      }	t        j                  d|	«      }	|	j                  |||z  |z  | j                  d   | j                  d   z  |z  f¬«      }	|	S )zJ
        pixel_values: [batch_size, num_frames, 3, height, width]
        rÊ   r   r   r   ©r>   zntchpwq->nthwpqc)r>   rš   rª   r)   Úeinsum)
ry   rA   rD   rt   r~   r   r€   Únum_patches_heightÚnum_patches_widthÚpatchified_pixel_valuess
             r/   Úpatchify_pixelz!TvltForPreTraining.patchify_pixel­  s  € ð ?K×>PÑ>PÑ;ˆ
�J ¨f°eØ)×/Ñ/°Ñ2°d×6KÑ6KÈAÑ6NÑNÐØ(×.Ñ.¨qÑ1°T×5JÑ5JÈ1Ñ5MÑMÐØ".×"6Ñ"6àØØØ"Ø×%Ñ% aÑ(Ø!Ø×%Ñ% aÑ(ðð #7ó 
#
Ðô #(§,¡,Ð/AÐCZÓ"[ÐØ"9×"AÑ"AàØ"Ð%6Ñ6¸ÑCØ×%Ñ% aÑ(¨4×+@Ñ+@ÀÑ+CÑCÀlÑRðð #Bó #
Ðð 'Ð&r.   c           	      óp  — |j                   \  }}}}|| j                  d   z  }|| j                  d   z  }|j                  |||| j                  d   || j                  d   f¬«      }t        j                  d|«      }|j                  |||z  | j                  d   | j                  d   z  |z  f¬«      }|S )z>
        audio_values: [batch_size, 1, height, width]
        r   r   r›  znchpwq->nhwpqc)r>   rŽ   rª   r)   rœ  )	ry   rO   rD   r~   r   r€   r�  rž  Úpatchified_audio_valuess	            r/   Úpatchify_audioz!TvltForPreTraining.patchify_audioÉ  sï   € ð 3?×2DÑ2DÑ/ˆ
�L &¨%Ø# t×'<Ñ'<¸QÑ'?Ñ?ÐØ! T×%:Ñ%:¸1Ñ%=Ñ=ÐØ".×"6Ñ"6àØØ"Ø×%Ñ% aÑ(Ø!Ø×%Ñ% aÑ(ðð #7ó 	#
Ðô #(§,¡,Ð/?ÐAXÓ"YÐØ"9×"AÑ"AàØ"Ð%6Ñ6Ø×%Ñ% aÑ(¨4×+@Ñ+@ÀÑ+CÑCÀlÑRðð #Bó #
Ðð 'Ð&r.   c                 ó¤   — | j                  |«      }||z
  dz  }|j                  d¬«      }||z  j                  «       |j                  «       z  }|S ©Nr;   rJ   rV   )r   r>  Úsum)ry   rA   Úpixel_predictionsÚmaskrŸ  r5   s         r/   Úpixel_mae_lossz!TvltForPreTraining.pixel_mae_lossä  óU   € Ø"&×"5Ñ"5°lÓ"CÐØ!Ð$;Ñ;ÀÑAˆØ�y‰y˜RˆyÓ ˆØ�t‘× Ñ Ó" T§X¡X£ZÑ/ˆØˆr.   c                 ó¤   — | j                  |«      }||z
  dz  }|j                  d¬«      }||z  j                  «       |j                  «       z  }|S r¥  )r£  r>  r¦  )ry   rO   Úaudio_predictionsr¨  r¢  r5   s         r/   Úaudio_mae_lossz!TvltForPreTraining.audio_mae_lossë  rª  r.   c           	      ó  — |j                   \  }}}|j                  ||j                   d   |z
  d«      }t        j                  ||gd¬«      }t        j                  |d|j                  d«      j                  dd|«      ¬«      }|S )Nr   rV   rJ   rX   )r>   rM   r)   re  r[   rL   )	ry   Ú
mask_tokenr]   ra   rD   Ú
seq_lengthrW   Úmask_tokensÚpadded_sequences	            r/   Úconcatenate_maskz#TvltForPreTraining.concatenate_maskò  s„   € Ø&.§n¡nÑ#ˆ
�J Ø ×'Ñ'¨
°K×4EÑ4EÀaÑ4HÈ:Ñ4UÐWXÓYˆÜŸ)™) X¨{Ð$;ÀÔCˆÜŸ,™,Ø ¨+×*?Ñ*?ÀÓ*C×*JÑ*JÈ1ÈaÐQTÓ*Uô
ˆð Ðr.   r\  rA   rO   rB   rP   ÚlabelsÚpixel_values_mixedÚpixel_mask_mixedrØ   r0  r1  r£   c                 óÔ  — |
�|
n| j                   j                  }
d}| j                  r~|€t        d«      ‚|€t        d«      ‚| j	                  ||||||	|
¬«      }|d   }| j                  |«      }t        «       } ||j                  d«      |j                  d«      «      }||z  }d}d}| j                  �rv| j                  �ri| j	                  ||||dd||	|
¬	«	      }|
r|j                  n|d
   }|
r|j                  n|d   }|
r|j                  n|d   }|
r|j                  n|d   }|
r|j                  n|d   }|
r|j                  n|d   }| j!                  |«      }| j!                  |«      }|j#                  d
«      }| j%                  | j&                  ||«      }|| j(                  j+                  d
|d
«      z   }|t-        j.                  | j0                  dd…d|…f   | j2                  d
¬«      z   }|| j4                  z   }| j7                  |«      }| j9                  |j:                  «      }| j%                  | j<                  ||«      }|j#                  d
«      | j>                  z  }|| j@                  j+                  d
|d
«      z   }|t-        j.                  | jB                  dd…d|…f   | j>                  d
¬«      z   }|| jD                  z   }| j7                  |«      }| jG                  |j:                  «      }| jI                  |||«      | jK                  |||«      z   }||z  }|
s||fdd z   }�|f|z   S |S tM        |||jN                  |jP                  ¬«      S )aF  
        pixel_values_mixed (`torch.FloatTensor` of shape `(batch_size, num_frames, num_channels, height, width)`):
            Pixel values that mix positive and negative samples in Tvlt vision-audio matching. Audio values can be
            obtained using [`TvltProcessor`]. See [`TvltProcessor.__call__`] for details.

        pixel_mask_mixed (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel masks of pixel_values_mixed. Pixel values mixed can be obtained using [`TvltProcessor`]. See
            [`TvltProcessor.__call__`] for details.

        labels (`torch.LongTensor` of shape `(batch_size, num_labels)`, *optional*):
            Labels for computing the vision audio matching loss. Indices should be in `[0, 1]`. num_labels has to be 1.

        Return:

        Examples:

        ```python
        >>> from transformers import TvltProcessor, TvltForPreTraining
        >>> import numpy as np
        >>> import torch

        >>> num_frames = 8
        >>> images = list(np.random.randn(num_frames, 3, 224, 224))
        >>> images_mixed = list(np.random.randn(num_frames, 3, 224, 224))
        >>> audio = list(np.random.randn(10000))
        >>> processor = TvltProcessor.from_pretrained("ZinengTang/tvlt-base")
        >>> model = TvltForPreTraining.from_pretrained("ZinengTang/tvlt-base")
        >>> input_dict = processor(
        ...     images, audio, images_mixed, sampling_rate=44100, mask_pixel=True, mask_audio=True, return_tensors="pt"
        ... )

        >>> outputs = model(**input_dict)
        >>> loss = outputs.loss
        ```Nr=  zMatching task requires labelsz)Matching task requires pixel_values_mixed©rB   rP   rØ   r0  r1  r   rJ   T)rB   rP   r^  r_  rØ   r0  r1  r   r;   rÊ   r   é   é   rV   é   )r5   r6   r7   r8   r#   r$   ))rw   ra  r†  r©   r;  r‰  r	   rN   r‡  r,  r   r   r   r    r!   r"   rŠ  r“   r³  r‹  rŽ  rM   r)   r}   r�  ro   r�  r�  r•  r2   rŒ  r�   r’  r‘  r“  r–  r©  r­  r4   r#   r$   ) ry   rA   rO   rB   rP   r´  rµ  r¶  rØ   r0  r1  Ú
total_lossrá   rr  r6   Úloss_fctr5   r7   r8   rs  rt  r   r    r!   r"   Úpixel_decoder_inputÚaudio_decoder_inputrt   Úpixel_decoder_outputsrS   Úaudio_decoder_outputsrð   s                                    r/   r‚   zTvltForPreTraining.forwardû  sæ  € ðb &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØˆ
à×ÒØˆ~Ü Ð!@ÓAÐAØ!Ð)Ü Ð!LÓMÐMà—i‘iØ"ØØ+Ø%Ø"3Ø%9Ø'ð  ó ˆGð & a™jˆOØ"×0Ñ0°ÓAˆOä(Ó*ˆHÙ˜O×0Ñ0°Ó4°f·k±kÀ"³oÓFˆDØ˜$ÑˆJàˆØˆØ�=‹=˜TŸ]›]Ø—i‘iØØØ%Ø%ØØØ"3Ø%9Ø'ð  ó 
ˆGñ HS G×$CÒ$CÐX_Ð`aÑXbÐ!ÙGR G×$CÒ$CÐX_Ð`aÑXbÐ!Ù=H × 9Ò 9ÈgÐVWÉjÐÙ=H × 9Ò 9ÈgÐVWÉjÐÙ=H × 9Ò 9ÈgÐVWÉjÐÙ=H × 9Ò 9ÈgÐVWÉjÐà"&×"9Ñ"9Ø%ó#Ðð #'×"9Ñ"9Ø%ó#Ðð &×*Ñ*¨1Ó-ˆJØ"&×"7Ñ"7¸×8MÑ8MÐObÐduÓ"vÐØ"5¸×8TÑ8T×8[Ñ8[Ð\]Ð_iÐklÓ8mÑ"mÐØ"5¼×8OÑ8OØ×+Ñ+ªA¨{°
¨{¨NÑ;¸T×=WÑ=WÐ]^ô9ñ #Ðð #6¸×8UÑ8UÑ"UÐØ$(§L¡LÐ1DÓ$EÐ!Ø×.Ñ.Ð/D×/KÑ/KÓLˆLà"&×"7Ñ"7¸×8MÑ8MÐObÐduÓ"vÐØ2×7Ñ7¸Ó:¸d×>SÑ>SÑSÐØ"5¸×8OÑ8O×8VÑ8VÐWXÐZjÐlmÓ8nÑ"nÐØ"5¼×8OÑ8OØ×,Ñ,ªQÐ0AÐ1AÐ0AÐ-AÑBÀD×DYÑDYÐ_`ô9ñ #Ðð #6¸×8UÑ8UÑ"UÐØ$(§L¡LÐ1DÓ$EÐ!Ø×.Ñ.Ð/D×/KÑ/KÓLˆLà×&Ñ& |°\ÐCTÓUÐX\×XkÑXkØ˜lÐ,=óYñ ˆDð ˜$ÑˆJáØ% |°\ÐBÀWÈQÈRÀ[ÑPˆFØ/3Ð/?�Z�M FÑ*ÐKÀVÐKä'ØØ+Ø%Ø%Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r.   )NNNNNNNN)r%   r&   r'   rl   r   r£  r©  r­  r³  r   ru  r   r4   rv  r)   r*   r   r,   rw  r   r   r‚   r…   r†   s   @r/   r„  r„  r  sG  ø„ ô
4òl'ò8'ò6òòñ +Ð+@ÓAÙÐ+CÐRaÔbð
 37Ø26Ø-1Ø:>Ø8<Ø,0Ø/3Ø&*ñH
à×'Ñ'ðH
ð ×'Ñ'ðH
ð ˜U×.Ñ.Ñ/ð	H
ð
 ˜U×.Ñ.Ñ/ðH
ð ˜×)Ñ)Ñ*ðH
ð % U×%6Ñ%6Ñ7ðH
ð # 5×#4Ñ#4Ñ5ðH
ð $ D™>ðH
ð ' t™nðH
ð ˜d‘^ðH
ð 
ˆu�U×&Ñ&Ñ'Ð)AÐAÑ	BòH
ó có BôH
r.   r„  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )Ú
TvltPoolerc                 ó²   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  «       | _        y rƒ   )rk   rl   r   rÁ   rr   ræ   ÚTanhÚ
activationrx   s     €r/   rl   zTvltPooler.__init__‰  s9   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
ÜŸ'™'›)ˆ�r.   c                 ó\   — |d d …df   }| j                  |«      }| j                  |«      }|S )Nr   )ræ   rÆ  )ry   r#   Úfirst_token_tensorÚpooled_outputs       r/   r‚   zTvltPooler.forwardŽ  s4   € Ø*ª1¨a¨4Ñ0ÐØŸ
™
Ð#5Ó6ˆØŸ™¨Ó6ˆØÐr.   r8  r†   s   @r/   rÃ  rÃ  ˆ  s   ø„ ô$ö
r.   rÃ  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )rˆ  c                 óŒ   •— t         ‰| �  «        t        |«      | _        t	        j
                  |j                  d«      | _        y rj   )rk   rl   rÃ  Úpoolerr   rÁ   rr   Úfcrx   s     €r/   rl   zTvltMatchingHead.__init__–  s2   ø€ Ü‰ÑÔÜ  Ó(ˆŒÜ—)‘)˜F×.Ñ.°Ó2ˆ�r.   c                 óF   — | j                  | j                  |«      «      }|S rƒ   )rÍ  rÌ  r  s     r/   r‚   zTvltMatchingHead.forward›  s   € ØŸ™ §¡¨MÓ :Ó;ˆØÐr.   r8  r†   s   @r/   rˆ  rˆ  •  s   ø„ ô3ö
r.   rˆ  c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )r”  c                 óz   •— t         ‰| �  «        || _        t        j                  |j
                  |«      | _        y rƒ   )rk   rl   rw   r   rÁ   r{  r�  )ry   rw   Ú
output_dimrz   s      €r/   rl   zTvltMAEHead.__init__¡  s-   ø€ Ü‰ÑÔØˆŒÜ—y‘y ×!;Ñ!;¸ZÓHˆ�r.   c                 ó(   — | j                  |«      }|S rƒ   )r�  r  s     r/   r‚   zTvltMAEHead.forward¦  s   € ØŸ™ ]Ó3ˆØÐr.   rƒ   r8  r†   s   @r/   r”  r”     s   ø„ õIö
r.   r”  zå
    Tvlt Model transformer with a classifier head on top (an MLP on top of the final hidden state of the [CLS] token)
    for audiovisual classification tasks, e.g. CMU-MOSEI Sentiment Analysis and Audio to Video Retrieval.
    c                   ó4  ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 	 	 	 	 	 dde	j                  de	j                  dee	j                     dee	j                     dee   dee   d	ee   d
ee	j                     deee	j                     ef   fd„«       «       Zˆ xZS )Ú TvltForAudioVisualClassificationc           	      óÔ  •— t         ‰| �  |«       t        |«      | _        t	        j
                  t	        j                  |j                  |j                  dz  «      t	        j                  |j                  dz  |j                  ¬«      t	        j                  «       t	        j                  |j                  dz  |j                  «      «      | _        || _        | j                  «        y )Nr;   r  )rk   rl   rM  r;  r   Ú
SequentialrÁ   rr   r  r  ÚGELUÚ
num_labelsÚ
classifierrw   rU  rx   s     €r/   rl   z)TvltForAudioVisualClassification.__init__³  s©   ø€ Ü‰Ñ˜Ô ä˜fÓ%ˆŒ	ô Ÿ-™-Ü�I‰I�f×(Ñ(¨&×*<Ñ*<¸qÑ*@ÓAÜ�L‰L˜×+Ñ+¨aÑ/°V×5JÑ5JÔKÜ�G‰G‹IÜ�I‰I�f×(Ñ(¨1Ñ,¨f×.?Ñ.?Ó@ó	
ˆŒð ˆŒð 	�‰Õr.   r\  rA   rO   rB   rP   rØ   r0  r1  r´  r£   c	           	      óÊ  — |�|n| j                   j                  }| j                  |||||||¬«      }	|	d   dd…df   }
| j                  |
«      }d}|�Y| j                   j                  dk(  rt        «       } |||«      }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, num_labels)`, *optional*):
            Labels for computing the audiovisual loss. Indices should be in `[0, ..., num_classes-1]` where num_classes
            refers to the number of classes in audiovisual tasks.

        Return:

        Examples:
        ```python
        >>> from transformers import TvltProcessor, TvltForAudioVisualClassification
        >>> import numpy as np
        >>> import torch

        >>> num_frames = 8
        >>> images = list(np.random.randn(num_frames, 3, 224, 224))
        >>> audio = list(np.random.randn(10000))
        >>> processor = TvltProcessor.from_pretrained("ZinengTang/tvlt-base")
        >>> model = TvltForAudioVisualClassification.from_pretrained("ZinengTang/tvlt-base")
        >>> input_dict = processor(images, audio, sampling_rate=44100, return_tensors="pt")

        >>> outputs = model(**input_dict)
        >>> loss = outputs.loss
        ```Nr¸  r   Ú
regressionÚclassificationr   )r5   r2   r#   r$   )
rw   ra  r;  rÙ  Ú	loss_typer   r
   r   r#   r$   )ry   rA   rO   rB   rP   rØ   r0  r1  r´  rá   rr  r2   r5   r½  rð   s                  r/   r‚   z(TvltForAudioVisualClassification.forwardÄ  s  € ðH &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØØ!Ø!Ø/Ø!5Ø#ð ó 
ˆð " !™*¢Q¨ TÑ*ˆØ—‘ Ó1ˆàˆØÐØ�{‰{×$Ñ$¨Ò4Ü"›9�Ù ¨Ó/‘Ø—‘×&Ñ&Ð*:Ò:Ü+Ó-�Ù ¨Ó/�áØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä'ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r.   )NNNNNN)r%   r&   r'   rl   r   ru  r   r   rv  r)   r*   r   rw  r,   r   r   r‚   r…   r†   s   @r/   rÔ  rÔ  «  sû   ø„ ôñ" +Ð+@ÓAÙÐ+CÐRaÔbð
 37Ø26Ø,0Ø/3Ø&*Ø-1ñB
à×'Ñ'ðB
ð ×'Ñ'ðB
ð ˜U×.Ñ.Ñ/ð	B
ð
 ˜U×.Ñ.Ñ/ðB
ð $ D™>ðB
ð ' t™nðB
ð ˜d‘^ðB
ð ˜×)Ñ)Ñ*ðB
ð 
ˆu�U×&Ñ&Ñ'Ð)AÐAÑ	BòB
ó có BôB
r.   rÔ  )Nç      è?)NrÞ  rK   é   rƒ   )Fr(   Úcollections.abcr�   rÒ   Úcopyr   Údataclassesr   Útypingr   r   r   r)   Útorch.utils.checkpointr   Útorch.nnr	   r
   r   Úactivationsr   Úmodeling_outputsr   r   Úmodeling_utilsr   Úpytorch_utilsr   r   Úutilsr   r   r   r   r   Úconfiguration_tvltr   Ú
get_loggerr%   Úloggerrv  Ú_CHECKPOINT_FOR_DOCr   r1   r4   rH   rT   re   ÚModulerg   rˆ   rm   rŠ   r·   rä   rí   rü   r  r
  r  r:  ÚTVLT_START_DOCSTRINGru  rM  ry  r„  rÃ  rˆ  r”  rÔ  r-   r.   r/   ú<module>rñ     sb  ðñ ã Û Ý Ý !ß )Ñ )ã Û Ý ß AÑ Aå "ß JÝ .ß R÷õ õ +ð 
ˆ×	Ñ	˜HÓ	%€à€Ø,Ð ð ô%?�kó %?ó ð%?ðP ô?˜ó ?ó ð?ð, ô?˜{ó ?ó ð?óBóó$Fô:+˜"Ÿ)™)ô +ô6+˜"Ÿ)™)ô +ô:&˜rŸy™yô &ôR)˜rŸy™yô )ôX9˜Ÿ	™	ô 9ôx�R—Y‘Yô ô$�B—I‘Iô ôD�r—y‘yô ô �—‘ô ô#�—	‘	ô #ôL2
�"—)‘)ô 2
ôj*˜/ô *ð0	Ð ð*Ð ñZ ØdØóô`
Ð#ó `
ó	ð`
ôF7q�"—)‘)ô 7qñt ØZØóôO
Ð,ó O
ó	ðO
ôd
�—‘ô 
ô�r—y‘yô ô�"—)‘)ô ñ ðð óôV
Ð':ó V
óñV
r.   