Ë
    T^(h¿  ã                   óH  — d Z ddlZddlmZ ddlmZ ddlmZm	Z	m
Z
mZmZ ddlZddlZddl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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(m)Z) ddl*m+Z+  e%jX                  e-«      Z.dZ/dZ0e G d„ de"«      «       Z1e G d„ de"«      «       Z2d„ Z3 G d„ dejh                  «      Z5 G d„ dejh                  «      Z6	 dBdejh                  dejn                  dejn                  dejn                  d e	ejn                     d!e8d"e8fd#„Z9 G d$„ d%ejh                  «      Z: G d&„ d'ejh                  «      Z; G d(„ d)ejh                  «      Z< G d*„ d+ejh                  «      Z= G d,„ d-ejh                  «      Z> G d.„ d/ejh                  «      Z? G d0„ d1ejh                  «      Z@ G d2„ d3e«      ZAd4ZBd5ZC e#d6eB«       G d7„ d8eA«      «       ZD G d9„ d:ejh                  «      ZE e#d;eB«       G d<„ d=eA«      «       ZF e#d>eB«       G d?„ d@eA«      «       ZGg dA¢ZHy)Cz,PyTorch VideoMAE (masked autoencoder) model.é    N)Údeepcopy)Ú	dataclass)ÚCallableÚOptionalÚSetÚTupleÚUnion)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )ÚACT2FN)ÚBaseModelOutputÚImageClassifierOutput)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)Ú find_pruneable_heads_and_indicesÚprune_linear_layer)ÚModelOutputÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstrings)ÚIMAGENET_DEFAULT_MEANÚIMAGENET_DEFAULT_STDé   )ÚVideoMAEConfigr   zMCG-NJU/videomae-basec                   ó–   — e Zd ZU dZdZeej                     ed<   dZ	ee
ej                        ed<   dZee
ej                        ed<   y)ÚVideoMAEDecoderOutputaO  
    Class for VideoMAEDecoder'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 + 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Úhidden_statesÚ
attentions)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r!   r   ÚtorchÚFloatTensorÚ__annotations__r"   r   r#   © ó    úl/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/videomae/modeling_videomae.pyr    r    1   sR   … ñð  +/€FˆH�U×&Ñ&Ñ'Ó.Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r,   r    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ej                        ed<   dZeeej                        ed<   y)ÚVideoMAEForPreTrainingOutputa±  
    Class for VideoMAEForPreTraining's outputs, with potential hidden states and attentions.

    Args:
        loss (`torch.FloatTensor` of shape `(1,)`):
            Pixel reconstruction loss.
        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 + 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Úlossr!   r"   r#   )r$   r%   r&   r'   r0   r   r(   r)   r*   r!   r"   r   r#   r+   r,   r-   r/   r/   H   sg   … ñð$ )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø*.€FˆH�U×&Ñ&Ñ'Ó.Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r,   r/   c                 óh  ‡— ˆfd„}t        j                  t        | «      D �cg c]
  } ||«      ‘Œ c}«      }t        j                  |dd…ddd…f   «      |dd…ddd…f<   t        j                  |dd…ddd…f   «      |dd…ddd…f<   t        j                  |«      j                  d«      S c c}w )z Sinusoid position encoding tablec           
      ó€   •— t        ‰«      D �cg c]$  }| t        j                  dd|dz  z  ‰z  «      z  ‘Œ& c}S c c}w )Ni'  é   )ÚrangeÚnpÚpower)ÚpositionÚhid_jÚd_hids     €r-   Úget_position_angle_vecz;get_sinusoid_encoding_table.<locals>.get_position_angle_vech   s;   ø€ ÜRWÐX]ÓR^Ö_È�œ2Ÿ8™8 E¨1°¸±
Ñ+;¸eÑ+CÓDÓDÒ_Ð_ùÒ_s   �);Nr   r3   r   )r5   Úarrayr4   ÚsinÚcosr(   r)   Ú	unsqueeze)Ú
n_positionr9   r:   Úpos_iÚsinusoid_tables    `   r-   Úget_sinusoid_encoding_tablerB   d   sª   ø€ ô`ô —X‘XÌ%ÐPZÓJ[Ö\ÀÑ5°eÕ<Ò\Ó]€NÜ Ÿf™f ^²A°q°t¸!°t°GÑ%<Ó=€N’1�a�d˜�d�7ÑÜ Ÿf™f ^²A°q°t¸!°t°GÑ%<Ó=€N’1�a�d˜�d�7Ñä×Ñ˜^Ó,×6Ñ6°qÓ9Ð9ùò	 ]s   £B/c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )ÚVideoMAEEmbeddingsz7
    Construct the patch and position embeddings.

    c                 óÐ   •— t         ‰| �  «        t        |«      | _        | j                  j                  | _        t        | j                  |j                  «      | _        || _        y ©N)	ÚsuperÚ__init__ÚVideoMAEPatchEmbeddingsÚpatch_embeddingsÚnum_patchesrB   Úhidden_sizeÚposition_embeddingsÚconfig©ÚselfrN   Ú	__class__s     €r-   rH   zVideoMAEEmbeddings.__init__x   sR   ø€ Ü‰ÑÔä 7¸Ó ?ˆÔØ×0Ñ0×<Ñ<ˆÔä#>¸t×?OÑ?OÐQW×QcÑQcÓ#dˆÔ Øˆ�r,   c                 ó  — | j                  |«      }|| j                  j                  «       j                  |«      j	                  |j
                  d¬«      z   }|�)|j                  \  }}}||    }|j                  |d|«      }|S )NT©ÚdeviceÚcopyéÿÿÿÿ)rJ   rM   ÚdetachÚtype_asÚtorT   ÚshapeÚreshape)rP   Úpixel_valuesÚbool_masked_posÚ
embeddingsÚ
batch_sizeÚ_Únum_channelss          r-   ÚforwardzVideoMAEEmbeddings.forward�   s—   € à×*Ñ*¨<Ó8ˆ
ð   $×":Ñ":×"AÑ"AÓ"C×"KÑ"KÈJÓ"W×"ZÑ"ZØ×$Ñ$¨4ð #[ó #
ñ 
ˆ
ð
 Ð&Ø*4×*:Ñ*:Ñ'ˆJ˜˜<Ø# _Ð$4Ñ5ˆJØ#×+Ñ+¨J¸¸LÓIˆJàÐr,   ©r$   r%   r&   r'   rH   rb   Ú__classcell__©rQ   s   @r-   rD   rD   r   s   ø„ ñô
ör,   rD   c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )rI   aw  
    Video to Patch Embedding. This module turns a batch of videos of shape (batch_size, num_frames, num_channels,
    height, width) into a tensor of shape (batch_size, seq_len, hidden_size) to be consumed by a Transformer encoder.

    The seq_len (the number of patches) equals (number of frames // tubelet_size) * (height // patch_size) * (width //
    patch_size).

    c           	      óˆ  •— t         ‰	| �  «        |j                  }|j                  }|j                  }|j
                  }|j                  }|j                  }t        |t        j                  j                  «      r|n||f}t        |t        j                  j                  «      r|n||f}|| _        || _        t        |«      | _        |d   |d   z  |d   |d   z  z  || j                  z  z  }|| _        || _        t        j                  ||| j                  |d   |d   f| j                  |d   |d   f¬«      | _        y )Nr   r   )Úin_channelsÚout_channelsÚkernel_sizeÚstride)rG   rH   Ú
image_sizeÚ
patch_sizera   rL   Ú
num_framesÚtubelet_sizeÚ
isinstanceÚcollectionsÚabcÚIterableÚintrK   r
   ÚConv3dÚ
projection)
rP   rN   rl   rm   ra   rL   rn   ro   rK   rQ   s
            €r-   rH   z VideoMAEPatchEmbeddings.__init__�   s>  ø€ Ü‰ÑÔà×&Ñ&ˆ
Ø×&Ñ&ˆ
Ø×*Ñ*ˆØ×(Ñ(ˆØ×&Ñ&ˆ
Ø×*Ñ*ˆä#-¨j¼+¿/¹/×:RÑ:RÔ#S‘ZÐZdÐfpÐYqˆ
Ü#-¨j¼+¿/¹/×:RÑ:RÔ#S‘ZÐZdÐfpÐYqˆ
Ø$ˆŒØ$ˆŒÜ Ó-ˆÔà˜‰]˜j¨™mÑ+°
¸1±ÀÈAÁÑ0NÑOÐS]Ðae×arÑarÑSrÑsð 	ð )ˆÔØ&ˆÔÜŸ)™)Ø$Ø$Ø×*Ñ*¨J°q©M¸:Àa¹=ÐIØ×%Ñ% z°!¡}°jÀ±mÐDô	
ˆ�r,   c                 ó”  — |j                   \  }}}}}|| j                  k7  rt        d«      ‚|| j                  d   k7  s|| j                  d   k7  r2t        d|› d|› d| j                  d   › d| j                  d   › d�	«      ‚|j	                  dddd	d
«      }| j                  |«      j                  d«      j                  dd«      }|S )NzeMake sure that the channel dimension of the pixel values match with the one set in the configuration.r   r   zInput image size (Ú*z) doesn't match model (z).r3   r   é   )rZ   ra   Ú
ValueErrorrl   Úpermuterv   ÚflattenÚ	transpose)rP   r\   r_   rn   ra   ÚheightÚwidthr^   s           r-   rb   zVideoMAEPatchEmbeddings.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óð ð $×+Ñ+¨A¨q°!°Q¸Ó:ˆØ—_‘_ \Ó2×:Ñ:¸1Ó=×GÑGÈÈ1ÓMˆ
ØÐr,   rc   re   s   @r-   rI   rI   “   s   ø„ ñô
ö6r,   rI   ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingÚdropoutc                 óÀ  — t        j                  ||j                  dd«      «      |z  }t        j                  j                  |dt         j                  ¬«      j                  |j                  «      }t        j                  j                  ||| j                  ¬«      }|�||z  }t        j                  ||«      }	|	j                  dd«      j                  «       }	|	|fS )NrV   éþÿÿÿ)ÚdimÚdtype)ÚpÚtrainingr   r3   )r(   Úmatmulr}   r
   Ú
functionalÚsoftmaxÚfloat32rY   rŠ   r†   rŒ   Ú
contiguous)
r€   r�   r‚   rƒ   r„   r…   r†   ÚkwargsÚattn_weightsÚattn_outputs
             r-   Úeager_attention_forwardr•   É   sÀ   € ô —<‘<  s§}¡}°R¸Ó'<Ó=ÀÑG€Lô —=‘=×(Ñ(¨¸2ÄUÇ]Á]Ð(ÓS×VÑVÐW\×WbÑWbÓc€Lô —=‘=×(Ñ(¨¸È6Ï?É?Ð(Ó[€Lð Ð!Ø# nÑ4ˆä—,‘,˜|¨UÓ3€KØ×'Ñ'¨¨1Ó-×8Ñ8Ó:€Kà˜Ð$Ð$r,   c            
       óè   ‡ — e Zd Zdeddfˆ fd„Zdej                  dej                  fd„Z	 d
deej                     de	de
eej                  ej                  f   eej                     f   fd	„Zˆ xZS )ÚVideoMAESelfAttentionrN   ÚreturnNc                 ó  •— t         ‰| �  «        |j                  |j                  z  dk7  r2t	        |d«      s&t        d|j                  › d|j                  › d�«      ‚|| _        |j                  | _        t        |j                  |j                  z  «      | _        | j                  | j                  z  | _	        |j                  | _        | j                  dz  | _        d| _        t        j                  |j                  | j                  d¬«      | _        t        j                  |j                  | j                  d¬«      | _        t        j                  |j                  | j                  d¬«      | _        |j&                  rot        j(                  t+        j,                  | j                  «      «      | _        t        j(                  t+        j,                  | j                  «      «      | _        y d | _        d | _        y )	Nr   Úembedding_sizezThe hidden size z4 is not a multiple of the number of attention heads ú.g      à¿F©Úbias)rG   rH   rL   Únum_attention_headsÚhasattrrz   rN   rt   Úattention_head_sizeÚall_head_sizeÚattention_probs_dropout_probÚdropout_probr…   Ú	is_causalr
   ÚLinearr�   r‚   rƒ   Úqkv_biasÚ	Parameterr(   ÚzerosÚq_biasÚv_biasrO   s     €r-   rH   zVideoMAESelfAttention.__init__è   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ˆÔØ"×?Ñ?ˆÔØ×/Ñ/°Ñ5ˆŒØˆŒä—Y‘Y˜v×1Ñ1°4×3EÑ3EÈEÔRˆŒ
Ü—9‘9˜V×/Ñ/°×1CÑ1CÈ%ÔPˆŒÜ—Y‘Y˜v×1Ñ1°4×3EÑ3EÈEÔRˆŒ
à�?Š?ÜŸ,™,¤u§{¡{°4×3EÑ3EÓ'FÓGˆDŒKÜŸ,™,¤u§{¡{°4×3EÑ3EÓ'FÓGˆD�KàˆDŒKØˆD�Kr,   Úxc                 ó¤   — |j                  «       d d | j                  | j                  fz   }|j                  |«      }|j	                  dddd«      S )NrV   r   r3   r   r   )Úsizerž   r    Úviewr{   )rP   r«   Únew_x_shapes      r-   Útranspose_for_scoresz*VideoMAESelfAttention.transpose_for_scores  sL   € Ø—f‘f“h˜s �m t×'?Ñ'?À×AYÑAYÐ&ZÑZˆØ�F‰F�;ÓˆØ�y‰y˜˜A˜q !Ó$Ð$r,   Ú	head_maskÚoutput_attentionsc           
      ó  — | j                   �!t        j                  | j                  d¬«      nd }t        j
                  j                  || j                  j                  |¬«      }t        j
                  j                  || j                  j                  | j                  ¬«      }t        j
                  j                  || j                  j                  | j                   ¬«      }| j                  |«      }| j                  |«      }	| j                  |«      }
t        }| j                  j                  dk7  rN| j                  j                  dk(  r|rt        j!                  d«       nt"        | j                  j                     } || |
||	|| j$                  | j&                  | j(                  sdn| j*                  ¬«      \  }}|j-                  «       d d	 | j.                  fz   }|j1                  |«      }|r||f}|S |f}|S )
NF)Úrequires_grad)ÚinputÚweightr�   ÚeagerÚsdpazã`torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to eager attention. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.ç        )r¤   r…   r†   rˆ   )r©   r(   Ú
zeros_likerª   r
   rŽ   Úlinearr‚   r¶   rƒ   r�   r°   r•   rN   Ú_attn_implementationÚloggerÚwarning_oncer   r¤   r…   rŒ   r£   r­   r¡   r[   )rP   r"   r±   r²   Úk_biasÚkeysÚvaluesÚqueriesÚ	key_layerÚvalue_layerÚquery_layerÚattention_interfaceÚcontext_layerÚattention_probsÚnew_context_layer_shapeÚoutputss                   r-   rb   zVideoMAESelfAttention.forward  sÂ  € ð HLÇ{Á{ÐG^”×!Ñ! $§+¡+¸UÕCÐdhˆÜ�}‰}×#Ñ#¨-ÀÇÁÇÁÐV\Ð#Ó]ˆÜ—‘×%Ñ%¨MÀ$Ç*Á*×BSÑBSÐZ^×ZeÑZeÐ%ÓfˆÜ—-‘-×&Ñ&¨]À4Ç:Á:×CTÑCTÐ[_×[fÑ[fÐ&Ógˆà×-Ñ-¨dÓ3ˆ	Ø×/Ñ/°Ó7ˆØ×/Ñ/°Ó8ˆä(?ÐØ�;‰;×+Ñ+¨wÒ6Ø�{‰{×/Ñ/°6Ò9Ñ>OÜ×#Ñ#ðLõô
 '>¸d¿k¹k×>^Ñ>^Ñ&_Ð#á)<ØØØØØØ—n‘nØ—L‘LØ#Ÿ}š}‘C°$×2CÑ2Cô	*
Ñ&ˆ�ð #0×"4Ñ"4Ó"6°s¸Ð";¸t×?QÑ?QÐ>SÑ"SÐØ%×-Ñ-Ð.EÓFˆá6G�= /Ð2ˆàˆð O\ÐM]ˆàˆr,   ©NF)r$   r%   r&   r   rH   r(   ÚTensorr°   r   Úboolr	   r   rb   rd   re   s   @r-   r—   r—   ç   sƒ   ø„ ð˜~ð °$õ ð4% e§l¡lð %°u·|±|ó %ð bgñ&Ø(0°·±Ñ(>ð&ØZ^ð&à	ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷&r,   r—   c                   ó|   ‡ — e Zd ZdZdeddfˆ fd„Zdej                  dej                  dej                  fd„Zˆ xZ	S )	ÚVideoMAESelfOutputz¥
    The residual connection is defined in VideoMAELayer instead of here (as is the case with other models), due to the
    layernorm applied before each block.
    rN   r˜   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  |j                  «      | _        y rF   )	rG   rH   r
   r¥   rL   ÚdenseÚDropoutÚhidden_dropout_probr†   rO   s     €r-   rH   zVideoMAESelfOutput.__init__7  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r,   r"   Úinput_tensorc                 óJ   — | j                  |«      }| j                  |«      }|S rF   ©rÑ   r†   ©rP   r"   rÔ   s      r-   rb   zVideoMAESelfOutput.forward<  s$   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆàÐr,   )
r$   r%   r&   r'   r   rH   r(   rÌ   rb   rd   re   s   @r-   rÏ   rÏ   1  sD   ø„ ñð
>˜~ð >°$õ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r,   rÏ   c                   óà   ‡ — e Zd Zdeddfˆ fd„Zdee   ddfd„Z	 	 ddej                  de
ej                     d	edeeej                  ej                  f   eej                     f   fd
„Zˆ xZS )ÚVideoMAEAttentionrN   r˜   Nc                 ó€   •— t         ‰| �  «        t        |«      | _        t	        |«      | _        t        «       | _        y rF   )rG   rH   r—   Ú	attentionrÏ   ÚoutputÚsetÚpruned_headsrO   s     €r-   rH   zVideoMAEAttention.__init__E  s0   ø€ Ü‰ÑÔÜ.¨vÓ6ˆŒÜ(¨Ó0ˆŒÜ›EˆÕr,   Úheadsc                 ó>  — t        |«      dk(  ry t        || j                  j                  | j                  j                  | j
                  «      \  }}t        | j                  j                  |«      | j                  _        t        | j                  j                  |«      | j                  _        t        | j                  j                  |«      | j                  _	        t        | j                  j                  |d¬«      | j                  _        | j                  j                  t        |«      z
  | j                  _        | j                  j                  | j                  j                  z  | j                  _        | j
                  j                  |«      | _        y )Nr   r   ©r‰   )Úlenr   rÛ   rž   r    rÞ   r   r�   r‚   rƒ   rÜ   rÑ   r¡   Úunion)rP   rß   Úindexs      r-   Úprune_headszVideoMAEAttention.prune_headsK  s  € Üˆu‹:˜Š?ØÜ7Ø�4—>‘>×5Ñ5°t·~±~×7YÑ7YÐ[_×[lÑ[ló
‰ˆˆuô
  2°$·.±.×2FÑ2FÈÓNˆ�‰ÔÜ/°·±×0BÑ0BÀEÓJˆ�‰ÔÜ1°$·.±.×2FÑ2FÈÓNˆ�‰ÔÜ.¨t¯{©{×/@Ñ/@À%ÈQÔOˆ�‰Ôð .2¯^©^×-OÑ-OÔRUÐV[ÓR\Ñ-\ˆ�‰Ô*Ø'+§~¡~×'IÑ'IÈDÏNÉN×LnÑLnÑ'nˆ�‰Ô$Ø ×-Ñ-×3Ñ3°EÓ:ˆÕr,   r"   r±   r²   c                 óh   — | j                  |||«      }| j                  |d   |«      }|f|dd  z   }|S )Nr   r   )rÛ   rÜ   )rP   r"   r±   r²   Úself_outputsÚattention_outputrÊ   s          r-   rb   zVideoMAEAttention.forward]  sE   € ð —~‘~ m°YÐ@QÓRˆàŸ;™; |°A¡¸ÓFÐà#Ð%¨°Q°RÐ(8Ñ8ˆØˆr,   rË   )r$   r%   r&   r   rH   r   rt   rå   r(   rÌ   r   rÍ   r	   r   rb   rd   re   s   @r-   rÙ   rÙ   D  s’   ø„ ð"˜~ð "°$õ "ð;  S¡ð ;¨dó ;ð* -1Ø"'ñ	à—|‘|ðð ˜EŸL™LÑ)ðð  ð	ð
 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷r,   rÙ   c                   ó`   ‡ — e Zd Zdeddfˆ fd„Zdej                  dej                  fd„Zˆ xZS )ÚVideoMAEIntermediaterN   r˜   Nc                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y rF   )rG   rH   r
   r¥   rL   Úintermediate_sizerÑ   rp   Ú
hidden_actÚstrr   Úintermediate_act_fnrO   s     €r-   rH   zVideoMAEIntermediate.__init__m  s]   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3KÑ3KÓLˆŒ
Ü�f×'Ñ'¬Ô-Ü'-¨f×.?Ñ.?Ñ'@ˆDÕ$à'-×'8Ñ'8ˆDÕ$r,   r"   c                 óJ   — | j                  |«      }| j                  |«      }|S rF   )rÑ   rï   )rP   r"   s     r-   rb   zVideoMAEIntermediate.forwardu  s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆàÐr,   ©	r$   r%   r&   r   rH   r(   rÌ   rb   rd   re   s   @r-   rê   rê   l  s1   ø„ ð9˜~ð 9°$õ 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 )ÚVideoMAEOutputrN   r˜   Nc                 óÈ   •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j                  «      | _	        y rF   )
rG   rH   r
   r¥   rì   rL   rÑ   rÒ   rÓ   r†   rO   s     €r-   rH   zVideoMAEOutput.__init__~  sB   ø€ Ü‰ÑÔÜ—Y‘Y˜v×7Ñ7¸×9KÑ9KÓLˆŒ
Ü—z‘z &×"<Ñ"<Ó=ˆ�r,   r"   rÔ   c                 óT   — | j                  |«      }| j                  |«      }||z   }|S rF   rÖ   r×   s      r-   rb   zVideoMAEOutput.forwardƒ  s.   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆà%¨Ñ4ˆàÐr,   rñ   re   s   @r-   ró   ró   }  s?   ø„ ð>˜~ð >°$õ >ð
 U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r,   ró   c                   óÎ   ‡ — e Zd ZdZdeddfˆ fd„Z	 	 d
dej                  deej                     de	de
eej                  ej                  f   eej                     f   fd	„Zˆ xZS )ÚVideoMAELayerz?This corresponds to the Block class in the timm implementation.rN   r˜   Nc                 ór  •— t         ‰| �  «        |j                  | _        d| _        t	        |«      | _        t        |«      | _        t        |«      | _	        t        j                  |j                  |j                  ¬«      | _        t        j                  |j                  |j                  ¬«      | _        y )Nr   ©Úeps)rG   rH   Úchunk_size_feed_forwardÚseq_len_dimrÙ   rÛ   rê   Úintermediateró   rÜ   r
   Ú	LayerNormrL   Úlayer_norm_epsÚlayernorm_beforeÚlayernorm_afterrO   s     €r-   rH   zVideoMAELayer.__init__�  s‡   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ*¨6Ó2ˆŒÜ0°Ó8ˆÔÜ$ VÓ,ˆŒÜ "§¡¨V×-?Ñ-?ÀV×EZÑEZÔ [ˆÔÜ!Ÿ|™|¨F×,>Ñ,>ÀF×DYÑDYÔZˆÕr,   r"   r±   r²   c                 óÞ   — | j                  | j                  |«      ||¬«      }|d   }|dd  }||z   }| j                  |«      }| j                  |«      }| j	                  ||«      }|f|z   }|S )N)r²   r   r   )rÛ   r   r  rý   rÜ   )rP   r"   r±   r²   Úself_attention_outputsrè   rÊ   Úlayer_outputs           r-   rb   zVideoMAELayer.forwardš  s–   € ð "&§¡Ø×!Ñ! -Ó0ØØ/ð "0ó "
Ðð
 2°!Ñ4ÐØ(¨¨Ð,ˆð )¨=Ñ8ˆð ×+Ñ+¨MÓ:ˆØ×(Ñ(¨Ó6ˆð —{‘{ <°Ó?ˆà�/ GÑ+ˆàˆr,   rË   )r$   r%   r&   r'   r   rH   r(   rÌ   r   rÍ   r	   r   rb   rd   re   s   @r-   r÷   r÷   �  s�   ø„ ÙIð[˜~ð [°$õ [ð -1Ø"'ñ	à—|‘|ðð ˜EŸL™LÑ)ðð  ð	ð
 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0°%¸¿¹Ñ2EÐEÑ	F÷r,   r÷   c                   óŠ   ‡ — e Zd Zdeddfˆ fd„Z	 	 	 	 ddej                  deej                     deded	ede	e
ef   fd
„Zˆ xZS )ÚVideoMAEEncoderrN   r˜   Nc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w rË   )
rG   rH   rN   r
   Ú
ModuleListr4   Únum_hidden_layersr÷   ÚlayerÚgradient_checkpointing)rP   rN   r`   rQ   s      €r-   rH   zVideoMAEEncoder.__init__¹  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]Ä5È×IaÑIaÓCbÖ#c¸a¤M°&Õ$9Ò#cÓdˆŒ
Ø&+ˆÕ#ùò $ds   ½A#r"   r±   r²   Úoutput_hidden_statesÚreturn_dictc                 ót  — |rdnd }|rdnd }t        | j                  «      D ]h  \  }}	|r||fz   }|�||   nd }
| j                  r+| j                  r| j	                  |	j
                  ||
|«      }n
 |	||
|«      }|d   }|sŒ`||d   fz   }Œj |r||fz   }|st        d„ |||fD «       «      S t        |||¬«      S )Nr+   r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrF   r+   ©Ú.0Úvs     r-   ú	<genexpr>z*VideoMAEEncoder.forward.<locals>.<genexpr>ã  s   è ø€ Òm˜qÐ_`Ñ_lœÑmùó   ‚Š©Úlast_hidden_stater"   r#   )Ú	enumerater
  r  rŒ   Ú_gradient_checkpointing_funcÚ__call__Útupler   )rP   r"   r±   r²   r  r  Úall_hidden_statesÚall_self_attentionsÚiÚlayer_moduleÚlayer_head_maskÚlayer_outputss               r-   rb   zVideoMAEEncoder.forward¿  sÿ   € ñ #7™B¸DÐÙ$5™b¸4Ðä(¨¯©Ó4ò 	P‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à.7Ð.C˜i¨šlÈˆOà×*Ò*¨t¯}ª}Ø $× AÑ AØ ×)Ñ)Ø!Ø#Ø%ó	!‘ñ !-¨]¸OÐM^Ó _�à)¨!Ñ,ˆMâ Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð'	Pñ*  Ø 1°]Ð4DÑ DÐáÜÑm ]Ð4EÐGZÐ$[ÔmÓmÐmÜØ+Ø+Ø*ô
ð 	
r,   )NFFT)r$   r%   r&   r   rH   r(   rÌ   r   rÍ   r	   r  r   rb   rd   re   s   @r-   r  r  ¸  sz   ø„ ð,˜~ð ,°$õ ,ð -1Ø"'Ø%*Ø ñ)
à—|‘|ð)
ð ˜EŸL™LÑ)ð)
ð  ð	)
ð
 #ð)
ð ð)
ð 
ˆu�oÐ%Ñ	&÷)
r,   r  c                   ó.   — e Zd ZdZeZdZdZdZdZ	dZ
d„ Zy)ÚVideoMAEPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úvideomaer\   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 weightsr¹   )ÚmeanÚstdNg      ð?)rp   r
   r¥   ru   r¶   ÚdataÚnormal_rN   Úinitializer_ranger�   Úzero_rþ   Úfill_)rP   r€   s     r-   Ú_init_weightsz%VideoMAEPreTrainedModel._init_weightsø  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_checkpointingÚ_supports_sdpaÚ_supports_flash_attn_2r,  r+   r,   r-   r"  r"  ë  s/   „ ñð
 "€LØ"ÐØ$€OØ&*Ð#Ø€NØ!Ðó
*r,   r"  aJ  
    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 ([`VideoMAEConfig`]): 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 [`AutoImageProcessor`]. See
            [`VideoMAEImageProcessor.__call__`] for details.

        head_mask (`torch.FloatTensor` of shape `(num_heads,)` or `(num_layers, num_heads)`, *optional*):
            Mask to nullify selected heads of the self-attention modules. Mask values selected in `[0, 1]`:

            - 1 indicates the head is **not masked**,
            - 0 indicates the head is **masked**.

        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
zbThe bare VideoMAE 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ej                     deej                     dee   d	ee   d
ee   deee	f   fd„«       «       Zˆ xZS )ÚVideoMAEModelc                 ó  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        |j                  rd | _        n0t        j                  |j                  |j                  ¬«      | _        | j                  «        y )Nrù   )rG   rH   rN   rD   r^   r  ÚencoderÚuse_mean_poolingÚ	layernormr
   rþ   rL   rÿ   Ú	post_initrO   s     €r-   rH   zVideoMAEModel.__init__,  si   ø€ Ü‰Ñ˜Ô ØˆŒä,¨VÓ4ˆŒÜ& vÓ.ˆŒà×"Ò"Ø!ˆD�NäŸ\™\¨&×*<Ñ*<À&×BWÑBWÔXˆDŒNð 	�‰Õr,   c                 ó.   — | j                   j                  S rF   )r^   rJ   )rP   s    r-   Úget_input_embeddingsz"VideoMAEModel.get_input_embeddings;  s   € Ø�‰×/Ñ/Ð/r,   c                 ó˜   — |j                  «       D ]7  \  }}| j                  j                  |   j                  j	                  |«       Œ9 y)z�
        Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base
        class PreTrainedModel
        N)Úitemsr6  r
  rÛ   rå   )rP   Úheads_to_pruner
  rß   s       r-   Ú_prune_headszVideoMAEModel._prune_heads>  sE   € ð
 +×0Ñ0Ó2ò 	C‰LˆE�5Ø�L‰L×Ñ˜uÑ%×/Ñ/×;Ñ;¸EÕBñ	Cr,   ©Úoutput_typer-  r\   r]   r±   r²   r  r  r˜   c                 óØ  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }| j	                  || j                   j
                  «      }| j                  ||«      }| j                  |||||¬«      }|d   }	| j                  �| j                  |	«      }	|s	|	f|dd z   S t        |	|j                  |j                  ¬«      S )a�  
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0). Each video in the
            batch must have the same number of masked patches. If `None`, then all patches are considered. Sequence
            length is `(num_frames // tubelet_size) * (image_size // patch_size) ** 2`.

        Returns:

        Examples:

        ```python
        >>> import av
        >>> import numpy as np

        >>> from transformers import AutoImageProcessor, VideoMAEModel
        >>> from huggingface_hub import hf_hub_download

        >>> np.random.seed(0)


        >>> def read_video_pyav(container, indices):
        ...     '''
        ...     Decode the video with PyAV decoder.
        ...     Args:
        ...         container (`av.container.input.InputContainer`): PyAV container.
        ...         indices (`List[int]`): List of frame indices to decode.
        ...     Returns:
        ...         result (np.ndarray): np array of decoded frames of shape (num_frames, height, width, 3).
        ...     '''
        ...     frames = []
        ...     container.seek(0)
        ...     start_index = indices[0]
        ...     end_index = indices[-1]
        ...     for i, frame in enumerate(container.decode(video=0)):
        ...         if i > end_index:
        ...             break
        ...         if i >= start_index and i in indices:
        ...             frames.append(frame)
        ...     return np.stack([x.to_ndarray(format="rgb24") for x in frames])


        >>> def sample_frame_indices(clip_len, frame_sample_rate, seg_len):
        ...     '''
        ...     Sample a given number of frame indices from the video.
        ...     Args:
        ...         clip_len (`int`): Total number of frames to sample.
        ...         frame_sample_rate (`int`): Sample every n-th frame.
        ...         seg_len (`int`): Maximum allowed index of sample's last frame.
        ...     Returns:
        ...         indices (`List[int]`): List of sampled frame indices
        ...     '''
        ...     converted_len = int(clip_len * frame_sample_rate)
        ...     end_idx = np.random.randint(converted_len, seg_len)
        ...     start_idx = end_idx - converted_len
        ...     indices = np.linspace(start_idx, end_idx, num=clip_len)
        ...     indices = np.clip(indices, start_idx, end_idx - 1).astype(np.int64)
        ...     return indices


        >>> # video clip consists of 300 frames (10 seconds at 30 FPS)
        >>> file_path = hf_hub_download(
        ...     repo_id="nielsr/video-demo", filename="eating_spaghetti.mp4", repo_type="dataset"
        ... )
        >>> container = av.open(file_path)

        >>> # sample 16 frames
        >>> indices = sample_frame_indices(clip_len=16, frame_sample_rate=1, seg_len=container.streams.video[0].frames)
        >>> video = read_video_pyav(container, indices)

        >>> image_processor = AutoImageProcessor.from_pretrained("MCG-NJU/videomae-base")
        >>> model = VideoMAEModel.from_pretrained("MCG-NJU/videomae-base")

        >>> # prepare video for the model
        >>> inputs = image_processor(list(video), return_tensors="pt")

        >>> # forward pass
        >>> outputs = model(**inputs)
        >>> last_hidden_states = outputs.last_hidden_state
        >>> list(last_hidden_states.shape)
        [1, 1568, 768]
        ```N©r±   r²   r  r  r   r   r  )rN   r²   r  Úuse_return_dictÚget_head_maskr	  r^   r6  r8  r   r"   r#   )
rP   r\   r]   r±   r²   r  r  Úembedding_outputÚencoder_outputsÚsequence_outputs
             r-   rb   zVideoMAEModel.forwardF  s  € ðx 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆð ×&Ñ& y°$·+±+×2OÑ2OÓPˆ	àŸ?™?¨<¸ÓIÐàŸ,™,ØØØ/Ø!5Ø#ð 'ó 
ˆð *¨!Ñ,ˆØ�>‰>Ð%Ø"Ÿn™n¨_Ó=ˆOáØ#Ð%¨¸¸Ð(;Ñ;Ð;äØ-Ø)×7Ñ7Ø&×1Ñ1ô
ð 	
r,   )NNNNN)r$   r%   r&   rH   r;  r?  r   ÚVIDEOMAE_INPUTS_DOCSTRINGr   r   Ú_CONFIG_FOR_DOCr(   r)   r   Ú
BoolTensorrÌ   rÍ   r	   r   rb   rd   re   s   @r-   r4  r4  '  sÌ   ø„ ô
ò0òCñ +Ð+DÓEÙ¨?ÈÔYð 7;Ø,0Ø,0Ø/3Ø&*ñ{
à×'Ñ'ð{
ð " %×"2Ñ"2Ñ3ð{
ð ˜EŸL™LÑ)ð	{
ð
 $ D™>ð{
ð ' t™nð{
ð ˜d‘^ð{
ð 
ˆu�oÐ%Ñ	&ò{
ó Zó Fô{
r,   r4  c                   ó,   ‡ — e Zd Zˆ fd„Z	 	 	 dd„Zˆ xZS )ÚVideoMAEDecoderc                 ó„  •— t         ‰| �  «        |j                  |j                  z  |j                  dz  z  }t        |«      }|j                  |_        |j                  |_	        |j                  |_        |j                  |_        t        j                  t!        |j                  «      D �cg c]  }t#        |«      ‘Œ c}«      | _        t        j&                  |j                  «      | _        |dkD  r t        j*                  |j                  |«      nt        j,                  «       | _        d| _        || _        y c c}w )Nr3   r   F)rG   rH   ra   ro   rm   r   Údecoder_hidden_sizerL   Údecoder_num_hidden_layersr	  Údecoder_num_attention_headsrž   Údecoder_intermediate_sizerì   r
   r  r4   r÷   Údecoder_layersrþ   Únormr¥   ÚIdentityÚheadr  rN   )rP   rN   rK   Údecoder_num_labelsÚdecoder_configr`   rQ   s         €r-   rH   zVideoMAEDecoder.__init__Ç  s  ø€ Ü‰ÑÔà#×0Ñ0°6×3FÑ3FÑFÈ×IZÑIZÐ\]ÑI]Ñ]Ðä! &Ó)ˆØ%+×%?Ñ%?ˆÔ"Ø+1×+KÑ+KˆÔ(Ø-3×-OÑ-OˆÔ*Ø+1×+KÑ+KˆÔ(Ü Ÿm™mÜ49¸&×:ZÑ:ZÓ4[Ö\¨qŒ]˜>Õ*Ò\ó
ˆÔô —L‘L ×!;Ñ!;Ó<ˆŒ	àI[Ð^_ÒI_ŒB�I‰I�f×0Ñ0Ð2DÔEÔeg×epÑepÓerð 	Œ	ð ',ˆÔ#Øˆ�ùò ]s   Â.D=c                 óÊ  — |rdnd }|rdnd }t        | j                  «      D ]`  \  }}	|r||fz   }| j                  r+| j                  r| j	                  |	j
                  |d |«      }
n |	|d |¬«      }
|
d   }|sŒX||
d   fz   }Œb |r||fz   }|dkD  r|d d …| d …f   }| j                  |«      }| j                  |«      }|st        d„ |||fD «       «      S t        |||¬«      S )Nr+   )r±   r²   r   r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrF   r+   r  s     r-   r  z*VideoMAEDecoder.forward.<locals>.<genexpr>  s   è ø€ Òf˜qÐXYÑXeœÑfùr  )r!   r"   r#   )
r  rS  r  rŒ   r  r  rT  rV  r  r    )rP   r"   Úreturn_token_numr²   r  r  r  r  r  r  r   r!   s               r-   rb   zVideoMAEDecoder.forwardÝ  s(  € ñ #7™B¸DÐÙ$5™b¸4ÐÜ(¨×)<Ñ)<Ó=ò 	P‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à×*Ò*¨t¯}ª}Ø $× AÑ AØ ×)Ñ)Ø!ØØ%ó	!‘ñ !-¨]ÀdÐ^oÔ p�à)¨!Ñ,ˆMâ Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð#	Pñ&  Ø 1°]Ð4DÑ DÐà˜aÒØ)ª!Ð.>Ð->Ñ-?Ð*?Ñ@ˆMð Ÿ	™	 -Ó0ˆØ—‘˜=Ó)ˆáÜÑf VÐ->Ð@SÐ$TÔfÓfÐfÜ$¨FÐBSÐ`sÔtÐtr,   )FFT)r$   r%   r&   rH   rb   rd   re   s   @r-   rM  rM  Æ  s   ø„ ôð4  Ø"Ø÷*ur,   rM  zXThe VideoMAE Model transformer with the decoder on top for self-supervised pre-training.c                   óÚ   ‡ — 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   dee   dee   d	eeef   fd
„«       «       Zˆ xZS )ÚVideoMAEForPreTrainingc                 ó  •— t         ‰| �  |«       || _        t        |«      | _        t        j                  |j                  |j                  d¬«      | _	        t        j                  t        j                  dd|j                  «      «      | _        t        | j                  j                  j                   |j                  «      | _        t%        || j                  j                  j                   ¬«      | _        | j)                  «        y )NFrœ   r   )rK   )rG   rH   rN   r4  r#  r
   r¥   rL   rO  Úencoder_to_decoderr§   r(   r¨   Ú
mask_tokenrB   r^   rK   rM   rM  Údecoderr9  rO   s     €r-   rH   zVideoMAEForPreTraining.__init__  s¼   ø€ Ü‰Ñ˜Ô ØˆŒä% fÓ-ˆŒä"$§)¡)¨F×,>Ñ,>À×@ZÑ@ZÐafÔ"gˆÔÜŸ,™,¤u§{¡{°1°a¸×9SÑ9SÓ'TÓUˆŒÜ#>Ø�M‰M×$Ñ$×0Ñ0°&×2LÑ2Ló$
ˆÔ ô ' v¸4¿=¹=×;SÑ;S×;_Ñ;_Ô`ˆŒð 	�‰Õr,   r@  r\   r]   r±   r²   r  r  r˜   c                 ó~  — |�|n| j                   j                  }| j                  ||||||¬«      }|d   }| j                  |«      }|j                  \  }	}
}|€t        d«      ‚| j                  j                  |	dd«      j                  |«      }|j                  «       j                  |j                  d¬«      }||    j                  |	d|«      }||   j                  |	d|«      }t        j                  ||z   | j                  |z   gd¬	«      }| j!                  ||j                  d   «      }|j"                  }d}t        j$                  «       5  | j                   j&                  d
k7  r|}nˆ|j                  }|j(                  }t        j*                  t,        «      j                  ||¬«      dddd…ddf   }t        j*                  t.        «      j                  ||¬«      dddd…ddf   }||z  |z   }|j                  \  }	}}}}| j                   j0                  | j                   j2                  }}| j                   j4                  rØ|j7                  |	||z  ||||z  |||z  |«      }|j9                  dddddddd
«      j;                  «       }|j7                  |	||z  |z  |z  |z  |z  ||z  |z  |«      }||j=                  dd¬«      z
  |j?                  ddd¬«      jA                  «       dz   z  }|j7                  |	||z  |z  |z  |z  |z  ||z  |z  |z  «      }n–| j                   j&                  d
k7  rt        d«      ‚|j7                  |	||z  ||||z  |||z  |«      }|j9                  dddddddd
«      j;                  «       }|j7                  |	||z  |z  |z  |z  |z  ||z  |z  |z  «      }|j                  \  }	}}||   j                  |	d|«      } ddd«       tC        «       }! |!| «      }|s|f|dd z   }"|�|f|"z   S |"S tE        |||jF                  |jH                  ¬«      S # 1 sw Y   ŒTxY w)a  
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, sequence_length)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0). Each video in the
            batch must have the same number of masked patches. Sequence length is `(num_frames // tubelet_size) *
            (image_size // patch_size) ** 2`.

        Returns:

        Examples:
        ```python
        >>> from transformers import AutoImageProcessor, VideoMAEForPreTraining
        >>> import numpy as np
        >>> import torch

        >>> num_frames = 16
        >>> video = list(np.random.randint(0, 256, (num_frames, 3, 224, 224)))

        >>> image_processor = AutoImageProcessor.from_pretrained("MCG-NJU/videomae-base")
        >>> model = VideoMAEForPreTraining.from_pretrained("MCG-NJU/videomae-base")

        >>> pixel_values = image_processor(video, return_tensors="pt").pixel_values

        >>> num_patches_per_frame = (model.config.image_size // model.config.patch_size) ** 2
        >>> seq_length = (num_frames // model.config.tubelet_size) * num_patches_per_frame
        >>> bool_masked_pos = torch.randint(0, 2, (1, seq_length)).bool()

        >>> outputs = model(pixel_values, bool_masked_pos=bool_masked_pos)
        >>> loss = outputs.loss
        ```N)r]   r±   r²   r  r  r   z!One must provided a boolean mask rV   TrS   r   rá   r   )rT   rŠ   ry   é   r3   é   é   rˆ   )r‰   Úkeepdim)r‰   Úunbiasedrf  g�íµ ÷Æ°>zQCan't unnormalize non-RGB images. Consider setting config.norm_pix_loss to False.©r0   r!   r"   r#   )%rN   rD  r#  r_  rZ   rz   rM   ÚexpandrX   rW   rY   rT   r[   r(   Úcatr`  ra  r!   Úno_gradra   rŠ   Ú	as_tensorr   r   ro   rm   Únorm_pix_lossr®   r{   r‘   r%  ÚvarÚsqrtr   r/   r"   r#   )#rP   r\   r]   r±   r²   r  r  rÊ   rH  r_   Úseq_lenra   Úexpanded_position_embeddingsÚpos_emb_visibleÚpos_emb_maskÚx_fullÚdecoder_outputsr!   r0   ÚframesrT   rŠ   r%  r&  Útimer~   r   ro   rm   Úframes_normÚvideos_patchr`   ÚlabelsÚloss_fctrÜ   s#                                      r-   rb   zVideoMAEForPreTraining.forward   sá  € ðP &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—-‘-ØØ+ØØ/Ø!5Ø#ð  ó 
ˆð " !™*ˆØ×1Ñ1Øó
ˆð -<×,AÑ,AÑ)ˆ
�G˜\ð Ð"ÜÐ@ÓAÐAØ'+×'?Ñ'?×'FÑ'FÀzÐSUÐWYÓ'Z×'bÑ'bÐcoÓ'pÐ$Ø'C×'JÑ'JÓ'L×'OÑ'OÐWc×WjÑWjÐquÐ'OÓ'vÐ$Ø6¸Ð7GÑH×PÑPÐQ[Ð]_ÐamÓnˆØ3°OÑD×LÑLÈZÐY[Ð]iÓjˆô —‘˜O¨oÑ=¸t¿¹ÐQ]Ñ?]Ð^ÐdeÔfˆð Ÿ,™, v¨|×/AÑ/AÀ!Ñ/DÓEˆØ ×'Ñ'ˆàˆÜ�]‰]‹_ñ H	Yà�{‰{×'Ñ'¨1Ò,à%‘ð &×,Ñ,�Ø$×*Ñ*�Ü—‘Ô'<Ó=×@Ñ@ÈÐV[Ð@Ó\Ð]aÐcgÒijÐlpÐrvÐ]vÑw�Ü—o‘oÔ&:Ó;×>Ñ>ÀfÐTYÐ>ÓZÐ[_ÐaeÒghÐjnÐptÐ[tÑu�Ø%¨Ñ+¨dÑ2�à<B¿L¹LÑ9ˆJ˜˜l¨F°EØ'+§{¡{×'?Ñ'?ÀÇÁ×AWÑAW˜*ˆLØ�{‰{×(Ò(àŸ™ØØ˜LÑ(Ø Ø Ø˜jÑ(ØØ˜ZÑ'Øó	�ð  Ÿ™¨¨1¨a°°A°q¸!¸QÓ?×JÑJÓL�àŸ™ØØ˜LÑ(¨6Ñ1°ZÑ?À%ÑGÈ:ÑUØ  :Ñ-°
Ñ:Ø ó	�ð  &¨¯©¸ÀD¨Ó(IÑIØ—J‘J 2°¸d�JÓC×HÑHÓJÈTÑQñ�ð  +×/Ñ/ØØ˜LÑ(¨6Ñ1°ZÑ?À%ÑGÈ:ÑUØ  :Ñ-°
Ñ:¸\ÑIó ‘ð —;‘;×+Ñ+¨qÒ0Ü$Økóð ð  Ÿ™ØØ˜LÑ(Ø Ø Ø˜jÑ(ØØ˜ZÑ'Øó	�ð  Ÿ™¨¨1¨a°°A°q¸!¸QÓ?×JÑJÓL�à%Ÿ{™{ØØ˜LÑ(¨6Ñ1°ZÑ?À%ÑGÈ:ÑUØ  :Ñ-°
Ñ:¸\ÑIó �ð +7×*<Ñ*<Ñ'ˆJ˜˜<Ø! /Ñ2×:Ñ:¸:ÀrÈ<ÓXˆF÷QH	YôT “9ˆÙ˜ Ó'ˆáØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä+ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
÷cH	Yð H	Yús   ÅJP3Ð3P<)NNNN)r$   r%   r&   rH   r   rI  r   r/   rJ  r(   r)   rK  r   rÌ   rÍ   r	   r  rb   rd   re   s   @r-   r]  r]  
  s¼   ø„ ô
ñ" +Ð+DÓEÙÐ+GÐVeÔfð
 -1Ø,0Ø/3Ø&*ñ]
à×'Ñ'ð]
ð ×)Ñ)ð]
ð ˜EŸL™LÑ)ð	]
ð
 $ D™>ð]
ð ' t™nð]
ð ˜d‘^ð]
ð 
ˆuÐ2Ð2Ñ	3ò]
ó gó Fô]
r,   r]  z£VideoMAE Model transformer with a video classification head on top (a linear layer on top of the average pooled hidden
    states of all tokens) e.g. for ImageNet.c                   óê   ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 	 	 	 	 	 dde	e
j                     de	e
j                     de	e
j                     de	e   de	e   de	e   d	eeef   fd
„«       «       Zˆ xZS )ÚVideoMAEForVideoClassificationc                 óŽ  •— t         ‰| �  |«       |j                  | _        t        |«      | _        |j
                  rt        j                  |j                  «      nd | _	        |j                  dkD  r*t        j                  |j                  |j                  «      nt        j                  «       | _        | j                  «        y )Nr   )rG   rH   Ú
num_labelsr4  r#  r7  r
   rþ   rL   Úfc_normr¥   rU  Ú
classifierr9  rO   s     €r-   rH   z'VideoMAEForVideoClassification.__init__È  s‘   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒÜ% fÓ-ˆŒð <B×;RÒ;R”r—|‘| F×$6Ñ$6Ô7ÐX\ˆŒØNT×N_ÑN_ÐbcÒNcœ"Ÿ)™) F×$6Ñ$6¸×8IÑ8IÔJÔik×itÑitÓivˆŒð 	�‰Õr,   r@  r\   r±   rz  r²   r  r  r˜   c                 ó‚  — |�|n| j                   j                  }| j                  |||||¬«      }|d   }| j                  �!| j                  |j	                  d«      «      }n	|dd…df   }| j                  |«      }	d}
|��‡| 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 )
a3  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the image classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).

        Returns:

        Examples:

        ```python
        >>> import av
        >>> import torch
        >>> import numpy as np

        >>> from transformers import AutoImageProcessor, VideoMAEForVideoClassification
        >>> from huggingface_hub import hf_hub_download

        >>> np.random.seed(0)


        >>> def read_video_pyav(container, indices):
        ...     '''
        ...     Decode the video with PyAV decoder.
        ...     Args:
        ...         container (`av.container.input.InputContainer`): PyAV container.
        ...         indices (`List[int]`): List of frame indices to decode.
        ...     Returns:
        ...         result (np.ndarray): np array of decoded frames of shape (num_frames, height, width, 3).
        ...     '''
        ...     frames = []
        ...     container.seek(0)
        ...     start_index = indices[0]
        ...     end_index = indices[-1]
        ...     for i, frame in enumerate(container.decode(video=0)):
        ...         if i > end_index:
        ...             break
        ...         if i >= start_index and i in indices:
        ...             frames.append(frame)
        ...     return np.stack([x.to_ndarray(format="rgb24") for x in frames])


        >>> def sample_frame_indices(clip_len, frame_sample_rate, seg_len):
        ...     '''
        ...     Sample a given number of frame indices from the video.
        ...     Args:
        ...         clip_len (`int`): Total number of frames to sample.
        ...         frame_sample_rate (`int`): Sample every n-th frame.
        ...         seg_len (`int`): Maximum allowed index of sample's last frame.
        ...     Returns:
        ...         indices (`List[int]`): List of sampled frame indices
        ...     '''
        ...     converted_len = int(clip_len * frame_sample_rate)
        ...     end_idx = np.random.randint(converted_len, seg_len)
        ...     start_idx = end_idx - converted_len
        ...     indices = np.linspace(start_idx, end_idx, num=clip_len)
        ...     indices = np.clip(indices, start_idx, end_idx - 1).astype(np.int64)
        ...     return indices


        >>> # video clip consists of 300 frames (10 seconds at 30 FPS)
        >>> file_path = hf_hub_download(
        ...     repo_id="nielsr/video-demo", filename="eating_spaghetti.mp4", repo_type="dataset"
        ... )
        >>> container = av.open(file_path)

        >>> # sample 16 frames
        >>> indices = sample_frame_indices(clip_len=16, frame_sample_rate=1, seg_len=container.streams.video[0].frames)
        >>> video = read_video_pyav(container, indices)

        >>> image_processor = AutoImageProcessor.from_pretrained("MCG-NJU/videomae-base-finetuned-kinetics")
        >>> model = VideoMAEForVideoClassification.from_pretrained("MCG-NJU/videomae-base-finetuned-kinetics")

        >>> inputs = image_processor(list(video), return_tensors="pt")

        >>> with torch.no_grad():
        ...     outputs = model(**inputs)
        ...     logits = outputs.logits

        >>> # model predicts one of the 400 Kinetics-400 classes
        >>> predicted_label = logits.argmax(-1).item()
        >>> print(model.config.id2label[predicted_label])
        eating spaghetti
        ```NrC  r   r   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationrV   rh  )rN   rD  r#  r€  r%  r�  Úproblem_typer  rŠ   r(   Úlongrt   r   Úsqueezer   r®   r   r   r"   r#   )rP   r\   r±   rz  r²   r  r  rÊ   rH  r!   r0   r{  rÜ   s                r-   rb   z&VideoMAEForVideoClassification.forwardÕ  s   € ð~ &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—-‘-ØØØ/Ø!5Ø#ð  ó 
ˆð " !™*ˆà�<‰<Ð#Ø"Ÿl™l¨?×+?Ñ+?ÀÓ+BÓC‰Oà-ªa°¨dÑ3ˆOà—‘ Ó1ˆàˆØÑØ�{‰{×'Ñ'Ð/Ø—?‘? 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,   )NNNNNN)r$   r%   r&   rH   r   rI  r   r   rJ  r   r(   rÌ   rÍ   r	   r   rb   rd   re   s   @r-   r}  r}  Â  sÇ   ø„ ôñ +Ð+DÓEÙÐ+@ÈÔ_ð 04Ø,0Ø)-Ø,0Ø/3Ø&*ñP
à˜uŸ|™|Ñ,ðP
ð ˜EŸL™LÑ)ðP
ð ˜Ÿ™Ñ&ð	P
ð
 $ D™>ðP
ð ' t™nðP
ð ˜d‘^ðP
ð 
ˆuÐ+Ð+Ñ	,òP
ó `ó FôP
r,   r}  )r]  r4  r"  r}  )r¹   )Ir'   Úcollections.abcrq   rU   r   Údataclassesr   Útypingr   r   r   r   r	   Únumpyr5   r(   Útorch.utils.checkpointr
   Útorch.nnr   r   r   Úactivationsr   Úmodeling_outputsr   r   Úmodeling_utilsr   r   Úpytorch_utilsr   r   Úutilsr   r   r   r   r   Úutils.constantsr   r   Úconfiguration_videomaer   Ú
get_loggerr$   r½   rJ  Ú_CHECKPOINT_FOR_DOCr    r/   rB   ÚModulerD   rI   rÌ   Úfloatr•   r—   rÏ   rÙ   rê   ró   r÷   r  r"  ÚVIDEOMAE_START_DOCSTRINGrI  r4  rM  r]  r}  Ú__all__r+   r,   r-   ú<module>rœ     s]  ðñ 3ã Ý Ý !ß 8Õ 8ã Û Û Ý ß AÑ Aå !ß Fß Fß Q÷õ ÷ KÝ 2ð 
ˆ×	Ñ	˜HÓ	%€à"€Ø-Ð ð ô:˜Kó :ó ð:ð, ô: ;ó :ó ð:ò6:ô˜Ÿ™ô ôB2˜bŸi™iô 2ðz ñ%Ø�I‰Ið%à�<‰<ð%ð 
�‰ð%ð �<‰<ð	%ð
 ˜UŸ\™\Ñ*ð%ð ð%ð ó%ô<F˜BŸI™Iô FôT˜Ÿ™ô ô&$˜Ÿ	™	ô $ôP˜2Ÿ9™9ô ô"�R—Y‘Yô ô '�B—I‘Iô 'ôV0
�b—i‘iô 0
ôf*˜oô *ð4	Ð ðÐ ñ. ØhØóôX
Ð+ó X
ó	ðX
ôvAu�b—i‘iô AuñH Ø^Øóôq
Ð4ó q
ó	ðq
ñh ð0àóô
`
Ð%<ó `
óð
`
òF s�r,   