§
    ‚Štj Ô  ã                   óV  — d dl Zd dlZd dl mZ d dlmZ d dlZd dlZd dl	m
c mZ d dlmZm
Z
 ddlm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 m!Z! ddl"m#Z# ddl$m%Z% ddl&m'Z'  e¦   «         rd dl(m)Z)  e!¦   «         rd dl*m+Z+ d dl,m-Z-  G d„ de
j.        ¦  «        Z/ G d„ de
j.        ¦  «        Z0 G d„ de
j.        ¦  «        Z1 G d„ de
j.        ¦  «        Z2	 dPde
j.        dej        dej        d ej        d!ej        dz  d"e3d#e3fd$„Z4 G d%„ d&e
j.        ¦  «        Z5 G d'„ d(e
j.        ¦  «        Z6 G d)„ d*e
j.        ¦  «        Z7 G d+„ d,e¦  «        Z8 G d-„ d.e
j.        ¦  «        Z9 e d/¬0¦  «        e G d1„ d2e¦  «        ¦   «         ¦   «         Z:	 dQd4ej        d5ej        d6ej        fd7„Z;d8ed9ed6efd:„Z<d8ej        d9ej        d6ej        fd;„Z= G d<„ d=e
j.        ¦  «        Z>d8ed9ed>e?d6efd?„Z@d8ej        d9ej        d>e?d6ej        fd@„ZA G dA„ dBe
j.        ¦  «        ZBe  G dC„ dDe¦  «        ¦   «         ZC G dE„ dFe
jD        ¦  «        ZE G dG„ dHe
j.        ¦  «        ZF G dI„ dJe
j.        ¦  «        ZG G dK„ dLe
j.        ¦  «        ZH e dM¬0¦  «         G dN„ dOeC¦  «        ¦   «         ZIdDdOgZJdS )Ré    N)ÚCallable)Ú	dataclass)ÚTensorÚnné   )Úinitialization)ÚACT2FN)ÚModelOutputÚis_scipy_availableÚrequires_backends)ÚGradientCheckpointingLayer)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚTransformersKwargsÚauto_docstringÚis_accelerate_available)Úmerge_with_config_defaults)Úcapture_outputsé   )ÚVideomtConfig)Úlinear_sum_assignment)ÚPartialState)Úreducec                   óF   ‡ — e Zd ZdZˆ fd„Zdej        dej        fd„Zˆ xZS )ÚVideomtPatchEmbeddingszì
    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        ||||¬¦  «        | _        d S )Nr   r   ©Úkernel_sizeÚstride)ÚsuperÚ__init__Ú
image_sizeÚ
patch_sizeÚnum_channelsÚhidden_sizeÚ
isinstanceÚcollectionsÚabcÚIterableÚnum_patchesr   ÚConv2dÚ
projection)ÚselfÚconfigr#   r$   r%   r&   r+   Ú	__class__s          €új/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/videomt/modeling_videomt.pyr"   zVideomtPatchEmbeddings.__init__:   sá   ø€ Ý‰Œ×ÒÑÔÐØ!'Ô!2°FÔ4E�Jˆ
Ø$*Ô$7¸Ô9K�kˆå#-¨j½+¼/Ô:RÑ#SÔ#SÐq�Z�ZÐZdÐfpÐYqˆ
Ý#-¨j½+¼/Ô:RÑ#SÔ#SÐq�Z�ZÐZdÐfpÐYqˆ
Ø! !”}¨
°1¬Ñ5¸*ÀQ¼-È:ÐVWÌ=Ñ:XÑYˆØ$ˆŒØ$ˆŒØ(ˆÔØ&ˆÔåœ) L°+È:Ð^hÐiÑiÔiˆŒˆˆó    Úpixel_valuesÚreturnc                 ó.  — |j         d         }|| j        k    rt          d| j        › d|› d�¦  «        ‚|                     | j        j        j        ¬¦  «        }|                      |¦  «                             d¦  «                             dd¦  «        }|S )Nr   zoMake sure that the channel dimension of the pixel values match with the one set in the configuration. Expected z	 but got ú.©Údtypeé   )	Úshaper%   Ú
ValueErrorÚtor-   Úweightr8   ÚflattenÚ	transpose)r.   r3   r%   Ú
embeddingss       r1   ÚforwardzVideomtPatchEmbeddings.forwardI   s®   € Ø#Ô)¨!Ô,ˆØ˜4Ô,Ò,Ð,ÝðIØ!Ô.ðIð IØ9EðIð Ið Iñô ð ð
 $—’¨T¬_Ô-CÔ-I�ÑJÔJˆØ—_’_ \Ñ2Ô2×:Ò:¸1Ñ=Ô=×GÒGÈÈ1ÑMÔMˆ
ØÐr2   )	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r"   Útorchr   rA   Ú__classcell__©r0   s   @r1   r   r   3   sm   ø€ € € € € ðð ðjð jð jð jð jð
 E¤Lð 
°U´\ð 
ð 
ð 
ð 
ð 
ð 
ð 
ð 
r2   r   c                   óf   ‡ — e Zd ZdZdeddfˆ fd„Zd	dej        dej        dz  dej        fd„Zˆ xZ	S )
ÚVideomtEmbeddingszM
    Construct the CLS token, mask token, position and patch embeddings.
    r/   r4   Nc                 óö  •— t          ¦   «                              ¦   «          || _        |j        | _        t	          j        t          j        dd|j        ¦  «        ¦  «        | _	        t	          j        t          j
        d|j        |j        ¦  «        ¦  «        | _        t          |¦  «        | _        | j        j        }t	          j        |j        ¦  «        | _        d|j        z   | _        t	          j        ||j        ¦  «        | _        |                      dt          j        |¦  «                             d¦  «        d¬¦  «         t	          j        t          j
        dd|j        ¦  «        ¦  «        | _        d S )Nr   Úposition_ids©r   éÿÿÿÿF)Ú
persistent)r!   r"   r/   r$   r   Ú	ParameterrF   Úrandnr&   Ú	cls_tokenÚzerosÚnum_register_tokensÚregister_tokensr   Úpatch_embeddingsr+   ÚDropoutÚhidden_dropout_probÚdropoutÚnum_prefix_tokensÚ	EmbeddingÚposition_embeddingsÚregister_bufferÚarangeÚexpandÚ
mask_token)r.   r/   r+   r0   s      €r1   r"   zVideomtEmbeddings.__init__[   s#  ø€ Ý‰Œ×ÒÑÔÐàˆŒØ Ô+ˆŒåœ¥e¤k°!°Q¸Ô8JÑ&KÔ&KÑLÔLˆŒÝ!œ|­E¬K¸¸6Ô;UÐW]ÔWiÑ,jÔ,jÑkÔkˆÔÝ 6°vÑ >Ô >ˆÔØÔ+Ô7ˆÝ”z &Ô"<Ñ=Ô=ˆŒØ!" VÔ%?Ñ!?ˆÔÝ#%¤<°¸VÔ=OÑ#PÔ#PˆÔ Ø×Ò˜^­U¬\¸+Ñ-FÔ-F×-MÒ-MÈgÑ-VÔ-VÐchÐÑiÔiÐiÝœ,¥u¤{°1°a¸Ô9KÑ'LÔ'LÑMÔMˆŒˆˆr2   r3   Úbool_masked_posc                 óö  — |j         dk    rD|j        \  }}}}}|                     ||z  |||¦  «        }|�|                     ||z  d¦  «        }n.|�,|j         dk    r!|                     |j        d         d¦  «        }|j        d         }|                      |¦  «        }|�T|                     |j        t          j        ¬¦  «                             d¦  «        }	t          j	        |	| j
        |¦  «        }| j                             |dd¦  «        }
| j                             |dd¦  «        }||                      | j        ¦  «        z   }t          j        |
||gd¬¦  «        }|                      |¦  «        }|S )Né   rN   r9   r   )Údevicer8   r   ©Údim)Úndimr:   ÚreshaperV   r<   rd   rF   ÚboolÚ	unsqueezeÚwherer`   rR   r_   rU   r\   rL   ÚcatrY   )r.   r3   ra   Ú
batch_sizeÚ
num_framesr%   ÚheightÚwidthr@   ÚmaskÚ
cls_tokensrU   s               r1   rA   zVideomtEmbeddings.forwardk   s}  € ØÔ Ò!Ð!ØBNÔBTÑ?ˆJ˜
 L°&¸%Ø'×/Ò/°
¸ZÑ0GÈÐW]Ð_dÑeÔeˆLàÐ*Ø"1×"9Ò"9¸*ÀzÑ:QÐSUÑ"VÔ"V�øØÐ(¨_Ô-AÀAÒ-EÐ-EØ-×5Ò5°oÔ6KÈAÔ6NÐPRÑSÔSˆOà!Ô'¨Ô*ˆ
Ø×*Ò*¨<Ñ8Ô8ˆ
àÐ&Ø"×%Ò%¨ZÔ->ÅeÄjÐ%ÑQÔQ×[Ò[Ð\^Ñ_Ô_ˆDÝœ T¨4¬?¸JÑGÔGˆJà”^×*Ò*¨:°r¸2Ñ>Ô>ˆ
ØÔ.×5Ò5°jÀ"ÀbÑIÔIˆà $×":Ò":¸4Ô;LÑ"MÔ"MÑMˆ
Ý”Y 
¨O¸ZÐHÈaÐPÑPÔPˆ
Ø—\’\ *Ñ-Ô-ˆ
ØÐr2   ©N©
rB   rC   rD   rE   r   r"   rF   r   rA   rG   rH   s   @r1   rJ   rJ   V   s™   ø€ € € € € ðð ðN˜}ð N°ð Nð Nð Nð Nð Nð Nð ð  E¤Lð À5Ä<ÐRVÑCVð ÐbgÔbnð ð ð ð ð ð ð ð r2   rJ   c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )Ú
VideomtMLPr4   Nc                 ó~  •— t          ¦   «                              ¦   «          |j        x}}t          |j        |j        z  ¦  «        }t          j        ||d¬¦  «        | _        t          |j	        t          ¦  «        rt          |j	                 | _        n|j	        | _        t          j        ||d¬¦  «        | _        d S )NT©Úbias)r!   r"   r&   ÚintÚ	mlp_ratior   ÚLinearÚfc1r'   Ú
hidden_actÚstrr	   Ú
activationÚfc2©r.   r/   Úin_featuresÚout_featuresÚhidden_featuresr0   s        €r1   r"   zVideomtMLP.__init__†   s¢   ø€ Ý‰Œ×ÒÑÔÐØ%+Ô%7Ð7ˆ�lÝ˜fÔ0°6Ô3CÑCÑDÔDˆÝ”9˜[¨/ÀÐEÑEÔEˆŒÝ�fÔ'­Ñ-Ô-ð 	0Ý$ VÔ%6Ô7ˆDŒOˆOà$Ô/ˆDŒOÝ”9˜_¨lÀÐFÑFÔFˆŒˆˆr2   Úhidden_statec                 ó„   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|S rs   )r}   r€   r�   ©r.   r†   s     r1   rA   zVideomtMLP.forward‘   s;   € Ø—x’x Ñ-Ô-ˆØ—’ |Ñ4Ô4ˆØ—x’x Ñ-Ô-ˆØÐr2   ©r4   N©rB   rC   rD   r"   rF   r   rA   rG   rH   s   @r1   rv   rv   …   si   ø€ € € € € ð	Gð 	Gð 	Gð 	Gð 	Gð 	Gð E¤Lð °U´\ð ð ð ð ð ð ð ð r2   rv   c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )ÚVideomtGatedMLPr4   Nc                 óD  •— t          ¦   «                              ¦   «          |j        x}}t          |j        |j        z  ¦  «        }t          |dz  dz  ¦  «        dz   dz  dz  }t          j        |d|z  d¬¦  «        | _        t          j        ||d¬¦  «        | _        d S ©Nr9   r   é   é   Trx   ©	r!   r"   r&   rz   r{   r   r|   Ú
weights_inÚweights_outr‚   s        €r1   r"   zVideomtGatedMLP.__init__™   ó    ø€ Ý‰Œ×ÒÑÔÐØ%+Ô%7Ð7ˆ�lÝ˜fÔ0°6Ô3CÑCÑDÔDˆÝ˜°Ñ2°QÑ6Ñ7Ô7¸!Ñ;ÀÑAÀAÑEˆåœ) K°°_Ñ1DÈ4ÐPÑPÔPˆŒÝœ9 _°lÈÐNÑNÔNˆÔÐÐr2   r†   c                 óÎ   — |                       |¦  «        }|                     dd¬¦  «        \  }}t          j                             |¦  «        |z  }|                      |¦  «        S ©Nr9   rN   re   ©r’   Úchunkr   Ú
functionalÚsilur“   ©r.   r†   Úx1Úx2Úhiddens        r1   rA   zVideomtGatedMLP.forward¢   ó]   € Ø—’ |Ñ4Ô4ˆØ×#Ò# A¨2Ð#Ñ.Ô.‰ˆˆBÝ”×#Ò# BÑ'Ô'¨"Ñ,ˆØ×Ò Ñ'Ô'Ð'r2   r‰   rŠ   rH   s   @r1   rŒ   rŒ   ˜   ói   ø€ € € € € ðOð Oð Oð Oð Oð Oð( E¤Lð (°U´\ð (ð (ð (ð (ð (ð (ð (ð (r2   rŒ   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingrY   c                 óÀ  — t          j        ||                     dd¦  «        ¦  «        |z  }|�||z   }t          j                             |dt           j        ¬¦  «                             |j        ¦  «        }t          j         	                    ||| j
        ¬¦  «        }t          j        ||¦  «        }	|	                     dd¦  «                             ¦   «         }	|	|fS )NrN   éþÿÿÿ)rf   r8   )ÚpÚtrainingr   r9   )rF   Úmatmulr?   r   r™   ÚsoftmaxÚfloat32r<   r8   rY   r«   Ú
contiguous)
r¢   r£   r¤   r¥   r¦   r§   rY   ÚkwargsÚattn_weightsÚattn_outputs
             r1   Úeager_attention_forwardr³   ©   sÃ   € õ ”<  s§}¢}°R¸Ñ'<Ô'<Ñ=Ô=ÀÑG€LØÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐW\ÔWbÑcÔc€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€Lå”,˜|¨UÑ3Ô3€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r2   c            
       ó~   ‡ — e Zd ZdZˆ fd„Z	 ddej        dej        dz  deej        ej        dz  f         fd„Zˆ xZ	S )	ÚVideomtAttentionz=Multi-headed attention from 'Attention Is All You Need' paperc                 ó‚  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        | j        | j        z  | _        | j        | j        z  | j        k    r t          d| j        › d| j        › d�¦  «        ‚| j        dz  | _	        |j
        | _        d| _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        d S )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: z).g      à¿F)r!   r"   r/   r&   Ú	embed_dimÚnum_attention_headsÚ	num_headsÚhead_dimr;   ÚscaleÚattention_dropoutrY   Ú	is_causalr   r|   Úk_projÚv_projÚq_projÚout_proj©r.   r/   r0   s     €r1   r"   zVideomtAttention.__init__Ã   s  ø€ Ý‰Œ×ÒÑÔÐØˆŒØÔ+ˆŒØÔ3ˆŒØœ¨$¬.Ñ8ˆŒØŒ=˜4œ>Ñ)¨T¬^Ò;Ð;Ýð'ÈdÌnð 'ð 'Ø”Nð'ð 'ð 'ñô ð ð ”] DÑ(ˆŒ
ØÔ/ˆŒØˆŒå”i ¤°´Ñ?Ô?ˆŒÝ”i ¤°´Ñ?Ô?ˆŒÝ”i ¤°´Ñ?Ô?ˆŒÝœ	 $¤.°$´.ÑAÔAˆŒˆˆr2   NÚhidden_statesr¦   r4   c           
      ó¼  — |j         dd…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }t          j        | j	        j
        t          ¦  «        }	 |	| ||||| j        | j        | j        sdn| j        ¬¦  «        \  }
} |
j        g |¢d‘R Ž                      ¦   «         }
|                      |
¦  «        }
|
|fS )z#Input shape: Batch x Time x ChannelNrN   r   r9   r¡   )r½   r§   rY   )r:   rº   rÀ   Úviewr?   r¾   r¿   r   Úget_interfacer/   Ú_attn_implementationr³   r½   r»   r«   rY   rh   r¯   rÁ   )r.   rÃ   r¦   r°   Úinput_shapeÚhidden_shapeÚqueriesÚkeysÚvaluesÚattention_interfacer²   r±   s               r1   rA   zVideomtAttention.forward×   sg  € ð $Ô)¨#¨2¨#Ô.ˆà8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆØ—+’+˜mÑ,Ô,×1Ò1°,Ñ?Ô?×IÒIÈ!ÈQÑOÔOˆØ�{Š{˜=Ñ)Ô)×.Ò.¨|Ñ<Ô<×FÒFÀqÈ!ÑLÔLˆØ—’˜]Ñ+Ô+×0Ò0°Ñ>Ô>×HÒHÈÈAÑNÔNˆå(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØØ”nØ”JØ#œ}Ð>�C�C°$´,ð	%
ñ 	%
ô 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—m’m KÑ0Ô0ˆà˜LÐ(Ð(r2   rs   )
rB   rC   rD   rE   r"   rF   r   ÚtuplerA   rG   rH   s   @r1   rµ   rµ   À   s™   ø€ € € € € ØGÐGðBð Bð Bð Bð Bð. /3ð!)ð !)à”|ð!)ð œ tÑ+ð!)ð
 
ˆuŒ|˜Uœ\¨DÑ0Ð0Ô	1ð!)ð !)ð !)ð !)ð !)ð !)ð !)ð !)r2   rµ   c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )ÚVideomtSwiGLUFFNr4   Nc                 óD  •— t          ¦   «                              ¦   «          |j        x}}t          |j        |j        z  ¦  «        }t          |dz  dz  ¦  «        dz   dz  dz  }t          j        |d|z  d¬¦  «        | _        t          j        ||d¬¦  «        | _        d S rŽ   r‘   r‚   s        €r1   r"   zVideomtSwiGLUFFN.__init__ü   r”   r2   r†   c                 óÎ   — |                       |¦  «        }|                     dd¬¦  «        \  }}t          j                             |¦  «        |z  }|                      |¦  «        S r–   r—   r›   s        r1   rA   zVideomtSwiGLUFFN.forward  rŸ   r2   r‰   rŠ   rH   s   @r1   rÐ   rÐ   û   r    r2   rÐ   c                   ó^   ‡ — e Zd ZdZd
deddfˆ fd„Zdej        dej        fd„Zde	fd	„Z
ˆ xZS )ÚVideomtDropPathzÏStochastic depth (DropPath) per sample, for residual blocks.

    Identity when ``drop_prob`` is 0 or outside training. See `Deep Networks with Stochastic Depth
    <https://arxiv.org/abs/1603.09382>`_.
    r¡   Ú	drop_probr4   Nc                 óV   •— t          ¦   «                              ¦   «          || _        d S rs   )r!   r"   rÕ   )r.   rÕ   r0   s     €r1   r"   zVideomtDropPath.__init__  s$   ø€ Ý‰Œ×ÒÑÔÐØ"ˆŒˆˆr2   rÃ   c                 ó  — | j         dk    s| j        s|S d| j         z
  }|j        d         fd|j        dz
  z  z   }t	          j        ||j        |j        ¬¦  «        }t	          j        ||z   ¦  «        }| 	                    |¦  «        |z  S )Nr¡   r   r   )r   ©r8   rd   )
rÕ   r«   r:   rg   rF   Úrandr8   rd   ÚfloorÚdiv)r.   rÃ   Ú	keep_probr:   Úrandom_tensors        r1   rA   zVideomtDropPath.forward  s“   € ØŒ>˜SÒ Ð ¨¬Ð Ø Ð Ø˜œÑ&ˆ	ØÔ$ QÔ'Ð)¨D°MÔ4FÈÑ4JÑ,KÑKˆÝœ
 5°Ô0CÈMÔL`ÐaÑaÔaˆÝœ M°IÑ$=Ñ>Ô>ˆØ× Ò  Ñ+Ô+¨mÑ;Ð;r2   c                 ó   — d| j         › �S )Nzp=)rÕ   ©r.   s    r1   Ú
extra_reprzVideomtDropPath.extra_repr   s   € Ø$�D”NÐ$Ð$Ð$r2   ©r¡   )rB   rC   rD   rE   Úfloatr"   rF   r   rA   r   rà   rG   rH   s   @r1   rÔ   rÔ     s›   ø€ € € € € ðð ð#ð # %ð #°$ð #ð #ð #ð #ð #ð #ð< U¤\ð <°e´lð <ð <ð <ð <ð%˜Cð %ð %ð %ð %ð %ð %ð %ð %r2   rÔ   c                   óh   ‡ — e Zd ZdZdeddfˆ fd„Z	 d	dej        dej        dz  dej        fd„Zˆ xZ	S )
ÚVideomtLayerzCThis corresponds to the Block class in the original implementation.r/   r4   Nc                 ó"  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¬¦  «        | _        t          |¦  «        | _        t          |¦  «        | _
        |j        dk    rt          |j        ¦  «        nt          j        ¦   «         | _        t          j        |j        |j        ¬¦  «        | _        |j        rt#          |¦  «        | _        nt'          |¦  «        | _        t          |¦  «        | _        d S )N©Úepsr¡   )r!   r"   r   Ú	LayerNormr&   Úlayer_norm_epsÚnorm1rµ   Ú	attentionÚVideomtLayerScaleÚlayer_scale1Údrop_path_raterÔ   ÚIdentityÚ	drop_pathÚnorm2Úuse_swiglu_ffnrÐ   Úmlprv   Úlayer_scale2rÂ   s     €r1   r"   zVideomtLayer.__init__'  sà   ø€ Ý‰Œ×ÒÑÔÐå”\ &Ô"4¸&Ô:OÐPÑPÔPˆŒ
Ý)¨&Ñ1Ô1ˆŒÝ-¨fÑ5Ô5ˆÔØCIÔCXÐ[^ÒC^ÐC^�¨Ô)>Ñ?Ô?Ð?ÕdfÔdoÑdqÔdqˆŒå”\ &Ô"4¸&Ô:OÐPÑPÔPˆŒ
àÔ ð 	*Ý'¨Ñ/Ô/ˆDŒHˆHå! &Ñ)Ô)ˆDŒHÝ-¨fÑ5Ô5ˆÔÐÐr2   rÃ   r¦   c                 ój  — |                       |¦  «        }|                      ||¦  «        \  }}|                      |¦  «        }|                      |¦  «        |z   }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        |z   }|S rs   )rê   rë   rí   rð   rñ   ró   rô   )r.   rÃ   r¦   Úhidden_states_normÚself_attention_outputÚ_Úlayer_outputs          r1   rA   zVideomtLayer.forward7  s³   € ð
 "ŸZšZ¨Ñ6Ô6ÐØ#'§>¢>Ð2DÀnÑ#UÔ#UÑ Ð˜qØ $× 1Ò 1Ð2GÑ HÔ HÐð ŸšÐ'<Ñ=Ô=ÀÑMˆð —z’z -Ñ0Ô0ˆØ—x’x Ñ-Ô-ˆØ×(Ò(¨Ñ6Ô6ˆð —~’~ lÑ3Ô3°mÑCˆàÐr2   rs   rt   rH   s   @r1   rä   rä   $  s–   ø€ € € € € ØMÐMð6˜}ð 6°ð 6ð 6ð 6ð 6ð 6ð 6ð& /3ðð à”|ðð œ tÑ+ðð 
Œð	ð ð ð ð ð ð ð r2   rä   c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )rì   r4   Nc                 ó¸   •— t          ¦   «                              ¦   «          t          j        |j        t          j        |j        ¦  «        z  ¦  «        | _        d S rs   )	r!   r"   r   rP   Úlayerscale_valuerF   Úonesr&   Úlambda1rÂ   s     €r1   r"   zVideomtLayerScale.__init__O  sC   ø€ Ý‰Œ×ÒÑÔÐÝ”| FÔ$;½e¼jÈÔI[Ñ>\Ô>\Ñ$\Ñ]Ô]ˆŒˆˆr2   r†   c                 ó   — || j         z  S rs   )rþ   rˆ   s     r1   rA   zVideomtLayerScale.forwardS  s   € Ø˜dœlÑ*Ð*r2   r‰   rŠ   rH   s   @r1   rì   rì   N  si   ø€ € € € € ð^ð ^ð ^ð ^ð ^ð ^ð+ E¤Lð +°U´\ð +ð +ð +ð +ð +ð +ð +ð +r2   rì   a¨  
    Class for outputs of [`VideomtForUniversalSegmentationOutput`].

    This output can be directly passed to [`~VideomtVideoProcessor.post_process_semantic_segmentation`] or
    [`~VideomtVideoProcessor.post_process_instance_segmentation`] or
    [`~VideomtVideoProcessor.post_process_panoptic_segmentation`] to compute final segmentation maps. Please, see
    [`~VideomtVideoProcessor`] for details regarding usage.
    )Úcustom_introc                   óà   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
ej        dz  ed<   dZeej                 dz  ed<   dZeej                 dz  ed<   dS )	Ú%VideomtForUniversalSegmentationOutputa€  
    loss (`torch.Tensor`, *optional*):
        The computed loss, returned when labels are present.
    class_queries_logits (`torch.FloatTensor`):
        A tensor of shape `(batch_size, num_queries, num_labels + 1)` representing the proposed classes for each
        query. Note the `+ 1` is needed because we incorporate the null class.
    masks_queries_logits (`torch.FloatTensor`):
        A tensor of shape `(batch_size, num_queries, height, width)` representing the proposed masks for each
        query.
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
        Last hidden states (final feature map) of the last layer.
    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 stage) of
        shape `(batch_size, sequence_length, hidden_size)`. Hidden-states all layers of the model.
    attentions (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `tuple(torch.FloatTensor)` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`. Self and Cross Attentions weights from transformer decoder.
    NÚlossÚclass_queries_logitsÚmasks_queries_logitsÚlast_hidden_staterÃ   Ú
attentions)rB   rC   rD   rE   r  rF   ÚFloatTensorÚ__annotations__r  r  r  rÃ   rÎ   r  © r2   r1   r  r  W  s»   € € € € € € ðð ð& &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø26Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø26€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6Ð6Ð6r2   r  FÚinput_featuresÚpoint_coordinatesr4   c                 óØ   — |                      ¦   «         dk    rd}|                     d¦  «        }t          j        j        j        | d|z  dz
  fi |¤Ž}|r|                     d¦  «        }|S )a(  
    A wrapper around `torch.nn.functional.grid_sample` to support 3D point_coordinates tensors.

    Args:
        input_features (`torch.Tensor` of shape (batch_size, channels, height, width)):
            A tensor that contains features map on a height * width grid
        point_coordinates (`torch.Tensor` of shape (batch_size, num_points, 2) or (batch_size, grid_height, grid_width,:
        2)):
            A tensor that contains [0, 1] * [0, 1] normalized point coordinates
        add_dim (`bool`):
            boolean value to keep track of added dimension

    Returns:
        point_features (`torch.Tensor` of shape (batch_size, channels, num_points) or (batch_size, channels,
        height_grid, width_grid):
            A tensor that contains features for points in `point_coordinates`.
    r   Tr9   g       @ç      ð?)rf   rj   rF   r   r™   Úgrid_sampleÚsqueeze)r  r  Úadd_dimr°   Úpoint_featuress        r1   Úsample_pointr    sƒ   € ð( ×ÒÑÔ !Ò#Ð#ØˆØ-×7Ò7¸Ñ:Ô:Ðõ ”XÔ(Ô4°^ÀSÐK\ÑE\Ð_bÑEbÐmÐmÐflÐmÐm€NØð 3Ø'×/Ò/°Ñ2Ô2ˆàÐr2   ÚinputsÚlabelsc                 ó(  — |                       ¦   «                              d¦  «        } dt          j        | |j        ¦  «        z  }|                      d¦  «        dd…df         |                     d¦  «        ddd…f         z   }d|dz   |dz   z  z
  }|S )aÈ  
    A pair wise version of the dice loss, see `dice_loss` for usage.

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        `torch.Tensor`: The computed loss between each pairs.
    r   r9   rN   N)Úsigmoidr>   rF   r¬   ÚTÚsum)r  r  Ú	numeratorÚdenominatorr  s        r1   Úpair_wise_dice_lossr  Ÿ  s�   € ð �^Š^ÑÔ×%Ò% aÑ(Ô(€FØ•E”L ¨¬Ñ2Ô2Ñ2€Ià—*’*˜R‘.”.    D Ô)¨F¯JªJ°r©N¬N¸4ÀÀÀ¸7Ô,CÑC€KØ�	˜A‘ +°¡/Ñ2Ñ2€DØ€Kr2   c                 óF  — | j         d         }t          j        d¬¦  «        } || t          j        | ¦  «        ¦  «        } || t          j        | ¦  «        ¦  «        }t          j        ||z  |j        ¦  «        }t          j        ||z  d|z
  j        ¦  «        }||z   }|S )aê  
    A pair wise version of the cross entropy loss, see `sigmoid_cross_entropy_loss` for usage.

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        loss (`torch.Tensor`): The computed loss between each pairs.
    r   Únone©Ú	reduction)r:   r   ÚBCEWithLogitsLossrF   Ú	ones_likeÚ
zeros_liker¬   r  )	r  r  Úheight_and_widthÚ	criterionÚcross_entropy_loss_posÚcross_entropy_loss_negÚloss_posÚloss_negr  s	            r1   Ú$pair_wise_sigmoid_cross_entropy_lossr*  µ  s£   € ð ”| A”ÐåÔ$¨vÐ6Ñ6Ô6€IØ&˜Y v­u¬¸vÑ/FÔ/FÑGÔGÐØ&˜Y v­uÔ/?ÀÑ/GÔ/GÑHÔHÐåŒ|Ð2Ð5EÑEÀvÄxÑPÔP€HÝŒ|Ð2Ð5EÑEÈÈFÉ
Ä~ÑVÔV€HØ�hÑ€DØ€Kr2   c                   ó¾   ‡ — e Zd ZdZ	 ddedededefˆ fd„Z ej        ¦   «         d	ej	        d
ej	        dej	        dej	        de
ee	                  f
d„¦   «         Zˆ xZS )ÚVideomtHungarianMatcheraq  This class computes an assignment between the labels and the predictions of the network.

    For efficiency reasons, the labels don't include the no_object. Because of this, in general, there are more
    predictions than labels. In this case, we do a 1-to-1 matching of the best predictions, while the others are
    un-matched (and thus treated as non-objects).
    r  é 1  Ú
cost_classÚ	cost_maskÚ	cost_diceÚ
num_pointsc                 óÂ   •— t          ¦   «                              ¦   «          |dk    r|dk    r|dk    rt          d¦  «        ‚|| _        || _        || _        || _        dS )aH  Creates the matcher

        Params:
            cost_class (`float`, *optional*, defaults to 1.0):
                Relative weight of the classification error in the matching cost.
            cost_mask (`float`, *optional*,  defaults to 1.0):
                This is the relative weight of the focal loss of the binary mask in the matching cost.
            cost_dice (`float`, *optional*, defaults to 1.0):
                This is the relative weight of the dice loss of the binary mask in the matching cost.
            num_points (`int`, *optional*, defaults to 12544):
                No. of points to sample on which the mask loss will be calculated. The same set of K points are
                uniformly sampled for all prediction and ground truth masks to construct the cost matrix for bipartite
                matching.
        r   zAll costs can't be 0N)r!   r"   r;   r1  r.  r/  r0  )r.   r.  r/  r0  r1  r0   s        €r1   r"   z VideomtHungarianMatcher.__init__Ù  sc   ø€ õ" 	‰Œ×ÒÑÔÐØ˜Š?ˆ?˜y¨Aš~˜~°)¸q².°.ÝÐ3Ñ4Ô4Ð4à$ˆŒØ$ˆŒØ"ˆŒØ"ˆŒˆˆr2   r  r  Úmask_labelsÚclass_labelsr4   c                 óH  — g }|j         d         }t          |¦  «        D �]õ}||                              d¦  «        }||         }	|dd…||         f          }
||                              |	¦  «        }|dd…df         }|	dd…df         }	t	          j        d| j        d|	j        ¬¦  «        }|                     |j         d         dd¦  «        }t          ||d¬¦  «         
                    d¦  «        }|                     |	j         d         dd¦  «        }t          |	|d¬¦  «         
                    d¦  «        }	t          |	|¦  «        }t          |	|¦  «        }| j        |z  | j        |
z  z   | j        |z  z   }t	          j        |t	          j        d	¦  «        ¦  «        }t	          j        |t	          j        d
¦  «        ¦  «        }t	          j        |d¦  «        }t)          |                     ¦   «         ¦  «        }|                     |¦  «         �Œ÷d„ |D ¦   «         }|S )ao  
        Params:
            masks_queries_logits (`torch.Tensor`):
                A tensor of dim `batch_size, num_queries, num_labels` with the classification logits.
            class_queries_logits (`torch.Tensor`):
                A tensor of dim `batch_size, num_queries, height, width` with the predicted masks.
            class_labels (`torch.Tensor`):
                A tensor of dim `num_target_boxes` (where num_target_boxes is the number of ground-truth objects in the
                target) containing the class labels.
            mask_labels (`torch.Tensor`):
                A tensor of dim `num_target_boxes, height, width` containing the target masks.

        Returns:
            matched_indices (`list[tuple[Tensor]]`): A list of size batch_size, containing tuples of (index_i, index_j)
            where:
                - index_i is the indices of the selected predictions (in order)
                - index_j is the indices of the corresponding selected labels (in order)
            For each batch element, it holds:
                len(index_i) = len(index_j) = min(num_queries, num_target_boxes).
        r   rN   Nr   r9   ©rd   F©Úalign_cornersg    _ Bg    _ Âc                 ó”   — g | ]E\  }}t          j        |t           j        ¬ ¦  «        t          j        |t           j        ¬ ¦  «        f‘ŒFS )r7   )rF   Ú	as_tensorÚint64)Ú.0ÚiÚjs      r1   ú
<listcomp>z3VideomtHungarianMatcher.forward.<locals>.<listcomp>5  sR   € ð 
ð 
ð 
Ù_cÐ_`Ðbc�UŒ_˜Q¥e¤kÐ2Ñ2Ô2µE´OÀAÍUÌ[Ð4YÑ4YÔ4YÐZð
ð 
ð 
r2   )r:   Úranger­   r<   rF   rÙ   r1  rd   Úrepeatr  r  r*  r  r/  r.  r0  ÚminimumÚtensorÚmaximumÚ
nan_to_numr   ÚcpuÚappend)r.   r  r  r3  r4  Úindicesrm   r=  Ú
pred_probsÚ	pred_maskr.  Útarget_maskr  Útarget_coordinatesÚpred_coordinatesr/  r0  Úcost_matrixÚassigned_indicesÚmatched_indicess                       r1   rA   zVideomtHungarianMatcher.forwardó  s6  € ð8 *,ˆð *Ô/°Ô2ˆ
Ý�zÑ"Ô"ð 	-ñ 	-ˆAØ-¨aÔ0×8Ò8¸Ñ<Ô<ˆJØ,¨QÔ/ˆIð % Q Q Q¨°Q¬Ð%7Ô8Ð8ˆJØ% aœ.×+Ò+¨IÑ6Ô6ˆKØ% a a a¨ gÔ.ˆKØ! ! ! ! T 'Ô*ˆIõ !&¤
¨1¨d¬o¸qÈÔIYÐ ZÑ ZÔ ZÐà!2×!9Ò!9¸+Ô:KÈAÔ:NÐPQÐSTÑ!UÔ!UÐÝ& {Ð4FÐV[Ð\Ñ\Ô\×dÒdÐefÑgÔgˆKà0×7Ò7¸	¼ÈÔ8JÈAÈqÑQÔQÐÝ$ YÐ0@ÐPUÐVÑVÔV×^Ò^Ð_`ÑaÔaˆIõ =¸YÈÑTÔTˆIå+¨I°{ÑCÔCˆIàœ.¨9Ñ4°t´ÈÑ7SÑSÐVZÔVdÐgpÑVpÑpˆKåœ-¨µU´\À$Ñ5GÔ5GÑHÔHˆKÝœ-¨µU´\À%Ñ5HÔ5HÑIÔIˆKÝÔ*¨;¸Ñ:Ô:ˆKå0EÀkÇoÂoÑFWÔFWÑ0XÔ0XÐØ�NŠNÐ+Ñ,Ô,Ð,Ñ,ð
ð 
Øgnð
ñ 
ô 
ˆð Ðr2   )r  r  r  r-  )rB   rC   rD   rE   râ   rz   r"   rF   Úno_gradr   ÚlistrÎ   rA   rG   rH   s   @r1   r,  r,  Ñ  sé   ø€ € € € € ðð ð joð#ð #Øð#Ø27ð#ØJOð#Øcfð#ð #ð #ð #ð #ð #ð4 €U„]�_„_ðDà#œlðDð $œlðDð ”\ð	Dð
 ”lðDð 
ˆe�FŒmÔ	ðDð Dð Dñ „_ðDð Dð Dð Dð Dr2   r,  Ú	num_masksc                 ó*  — |                       ¦   «                              d¦  «        }d||z                       d¦  «        z  }|                     d¦  «        |                     d¦  «        z   }d|dz   |dz   z  z
  }|                     ¦   «         |z  }|S )a4  
    Compute the DICE loss, similar to generalized IOU for masks as follows:

    $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x \cap y }{x \cup y + 1}} $$

    In practice, since `labels` is a binary mask, (only 0s and 1s), dice can be computed as follow

    $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x * y }{x + y + 1}} $$

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).
        num_masks (`int`):
            The number of masks present in the current batch, used for normalization.

    Returns:
        `torch.Tensor`: The computed loss.
    r   r9   rN   )r  r>   r  )r  r  rS  Úprobsr  r  r  s          r1   Ú	dice_lossrV  ;  s‰   € ð, �NŠNÑÔ×$Ò$ QÑ'Ô'€EØ�U˜V‘^×(Ò(¨Ñ,Ô,Ñ,€IØ—)’)˜B‘-”- &§*¢*¨R¡.¤.Ñ0€KØ�	˜A‘ +°¡/Ñ2Ñ2€DØ�8Š8‰:Œ:˜	Ñ!€DØ€Kr2   c                 óœ   — t          j        d¬¦  «        } || |¦  «        }|                     d¦  «                             ¦   «         |z  }|S )a|  
    Args:
        inputs (`torch.Tensor`):
            A float tensor of arbitrary shape.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        loss (`torch.Tensor`): The computed loss.
    r  r  r   )r   r!  Úmeanr  )r  r  rS  r%  Úcross_entropy_lossr  s         r1   Úsigmoid_cross_entropy_lossrZ  Y  sR   € õ Ô$¨vÐ6Ñ6Ô6€IØ"˜ 6¨6Ñ2Ô2Ðà×"Ò" 1Ñ%Ô%×)Ò)Ñ+Ô+¨iÑ7€DØ€Kr2   c                   ó~  ‡ — e Zd Zdedeeef         fˆ fd„Zdeee	                  dee	         fd„Z
dee         deeef         fd„Zd	ed
ee         deej                 deeef         fd„Zdej        deej                 deej                 de	deeej        f         f
d„Zd„ Zd„ Zdej        dej        fd„Zdej        de	de	dedej        f
d„Z	 ddej        d	ej        deej                 d
eej                 deeej        f         dz  deeej        f         fd„Zd
ej        dej        dej        fd„Zˆ xZS )ÚVideomtLossr/   Úweight_dictc                 óÀ  •— t          ¦   «                              ¦   «          t          | dg¦  «         |j        | _        || _        |j        | _        t          j        | j        dz   ¦  «        }| j        |d<   |  	                    d|¦  «         |j
        | _        |j        | _        |j        | _        t          |j        |j        |j        | j        ¬¦  «        | _        dS )aQ  
        The Videomt Loss. The loss is computed very similar to DETR. The process happens in two steps: 1) we
        compute hungarian assignment between ground truth masks and the outputs of the model 2) we supervise each pair
        of matched ground-truth / prediction (supervise class and mask)

        Args:
            config (`VideomtConfig`):
                The configuration for Videomt model also containing loss calculation specific parameters.
            weight_dict (`dict[str, float]`):
                A dictionary of weights to be applied to the different losses.
        Úscipyr   rN   Úempty_weight)r.  r0  r/  r1  N)r!   r"   r   Ú
num_labelsr]  Úno_object_weightÚeos_coefrF   rý   r]   Útrain_num_pointsr1  Úoversample_ratioÚimportance_sample_ratior,  Úclass_weightÚdice_weightÚmask_weightÚmatcher)r.   r/   r]  r`  r0   s       €r1   r"   zVideomtLoss.__init__n  sÖ   ø€ õ 	‰Œ×ÒÑÔÐÝ˜$  	Ñ*Ô*Ð*Ø Ô+ˆŒØ&ˆÔð Ô/ˆŒÝ”z $¤/°AÑ"5Ñ6Ô6ˆØœ=ˆ�RÑØ×Ò˜^¨\Ñ:Ô:Ð:ð !Ô1ˆŒØ &Ô 7ˆÔØ'-Ô'EˆÔ$å.ØÔ*ØÔ(ØÔ(Ø”ð	
ñ 
ô 
ˆŒˆˆr2   Úsizesr4   c                 óŒ   — |d         }|dd …         D ]0}t          |¦  «        D ]\  }}t          ||         |¦  «        ||<   ŒŒ1|S )Nr   r   )Ú	enumerateÚmax)r.   rk  ÚmaxesÚsublistÚindexÚitems         r1   Ú_max_by_axiszVideomtLoss._max_by_axis‘  s`   € Ø�a”ˆØ˜Q˜R˜R”yð 	7ð 	7ˆGÝ(¨Ñ1Ô1ð 7ð 7‘��tÝ" 5¨¤<°Ñ6Ô6��e‘�ð7àˆr2   Útensorsc                 ó"  — |                       d„ |D ¦   «         ¦  «        }t          |¦  «        g|z   }|\  }}}}|d         j        }|d         j        }	t	          j        |||	¬¦  «        }
t	          j        |||ft          j        |	¬¦  «        }t          ||
|¦  «        D ]l\  }}}|d |j	        d         …d |j	        d         …d |j	        d         …f          
                    |¦  «         d|d |j	        d         …d |j	        d         …f<   Œm|
|fS )Nc                 ó6   — g | ]}t          |j        ¦  «        ‘ŒS r
  )rR  r:   )r<  rC  s     r1   r?  z;VideomtLoss._pad_images_to_max_in_batch.<locals>.<listcomp>›  s"   € Ð%OÐ%OÐ%O¸V¥d¨6¬<Ñ&8Ô&8Ð%OÐ%OÐ%Or2   r   rØ   r   r9   F)rs  Úlenr8   rd   rF   rS   rý   ri   Úzipr:   Úcopy_)r.   rt  Úmax_sizeÚbatch_shaperm   rø   ro   rp   r8   rd   Úpadded_tensorsÚpadding_masksrC  Úpadded_tensorÚpadding_masks                  r1   Ú_pad_images_to_max_in_batchz'VideomtLoss._pad_images_to_max_in_batch™  s1  € à×$Ò$Ð%OÐ%OÀwÐ%OÑ%OÔ%OÑPÔPˆå˜7‘|”|�n xÑ/ˆØ'2Ñ$ˆ
�A�v˜uØ˜”
Ô ˆØ˜”Ô"ˆÝœ [¸ÀfÐMÑMÔMˆÝœ
 J°¸Ð#>ÅeÄjÐY_Ð`Ñ`Ô`ˆå36°wÀÐP]Ñ3^Ô3^ð 	Gð 	GÑ/ˆF�M <ØÐ+˜FœL¨œOÐ+Ð->¨v¬|¸A¬Ð->Ð@QÀ&Ä,ÈqÄ/Ð@QÐQÔR×XÒXÐY_Ñ`Ô`Ð`ØAFˆLÐ*˜6œ<¨œ?Ð*Ð,=¨f¬l¸1¬oÐ,=Ð=Ñ>Ð>à˜}Ð,Ð,r2   r  r4  rH  c                 óˆ  — |}|j         \  }}}t          j        | j        ¬¦  «        }|                      |¦  «        }	t          j        d„ t          ||¦  «        D ¦   «         ¦  «        }
t          j        ||f| j	        t
          j
        |j        ¬¦  «        }|
||	<   |                     dd¦  «        } |||¦  «        }d|i}|S )a…  Compute the losses related to the labels using cross entropy.

        Args:
            class_queries_logits (`torch.Tensor`):
                A tensor of shape `batch_size, num_queries, num_labels`
            class_labels (`list[torch.Tensor]`):
                List of class labels of shape `(labels)`.
            indices (`tuple[np.array])`:
                The indices computed by the Hungarian matcher.

        Returns:
            `dict[str, Tensor]`: A dict of `torch.Tensor` containing the following key:
            - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels.
        )r=   c                 ó*   — g | ]\  }\  }}||         ‘ŒS r
  r
  )r<  Útargetrø   r>  s       r1   r?  z+VideomtLoss.loss_labels.<locals>.<listcomp>À  s$   € ÐHÐHÐH™>˜6¡6 A qˆV�AŒYÐHÐHÐHr2   )Ú
fill_valuer8   rd   r   r9   Úloss_cross_entropy)r:   r   ÚCrossEntropyLossr`  Ú$_get_predictions_permutation_indicesrF   rl   rx  Úfullra  r;  rd   r?   )r.   r  r4  rH  Úpred_logitsrm   Únum_queriesrø   r%  ÚidxÚtarget_classes_oÚtarget_classesÚpred_logits_transposedÚloss_ceÚlossess                  r1   Úloss_labelszVideomtLoss.loss_labelsª  sß   € ð" +ˆØ%0Ô%6Ñ"ˆ
�K ÝÔ'¨tÔ/@ÐAÑAÔAˆ	Ø×7Ò7¸Ñ@Ô@ˆÝ œ9ØHÐH­S°¸wÑ-GÔ-GÐHÑHÔHñ
ô 
Ðõ œØ˜Ð%°$´/ÍÌÐ]hÔ]oð
ñ 
ô 
ˆð /ˆ�sÑà!,×!6Ò!6°q¸!Ñ!<Ô!<ÐØ�)Ð2°NÑCÔCˆØ&¨Ð0ˆØˆr2   r  r3  rS  c                 óf  ‡ — ‰                       |¦  «        }‰                      |¦  «        }||         }‰                      |¦  «        \  }}	||         }|dd…df         }|dd…df         }t          j        ¦   «         5  ‰                      |ˆ fd„‰ j        ‰ j        ‰ j        ¦  «        }
t          ||
d¬¦  «         
                    d¦  «        }ddd¦  «         n# 1 swxY w Y   t          ||
d¬¦  «         
                    d¦  «        }t          |||¦  «        t          |||¦  «        dœ}~~|S )a¤  Compute the losses related to the masks using sigmoid_cross_entropy_loss and dice loss.

        Args:
            masks_queries_logits (`torch.Tensor`):
                A tensor of shape `(batch_size, num_queries, height, width)`.
            mask_labels (`torch.Tensor`):
                List of mask labels of shape `(labels, height, width)`.
            indices (`tuple[np.array])`:
                The indices computed by the Hungarian matcher.
            num_masks (`int)`:
                The number of masks, used for normalization.

        Returns:
            losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing two keys:
            - **loss_mask** -- The loss computed using sigmoid cross entropy loss on the predicted and ground truth.
              masks.
            - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth,
              masks.
        Nc                 ó.   •— ‰                      | ¦  «        S rs   )Úcalculate_uncertainty)Úlogitsr.   s    €r1   ú<lambda>z(VideomtLoss.loss_masks.<locals>.<lambda>÷  s   ø€ ˜t×9Ò9¸&ÑAÔA€ r2   Fr7  r   )Ú	loss_maskÚ	loss_dice)r‡  Ú _get_targets_permutation_indicesr€  rF   rQ  Úsample_points_using_uncertaintyr1  re  rf  r  r  rZ  rV  )r.   r  r3  rH  rS  Úsrc_idxÚtgt_idxÚ
pred_masksÚtarget_masksrø   r  Úpoint_labelsÚpoint_logitsr�  s   `             r1   Ú
loss_maskszVideomtLoss.loss_masksÌ  s°  ø€ ð4 ×;Ò;¸GÑDÔDˆØ×7Ò7¸Ñ@Ô@ˆà)¨'Ô2ˆ
ð ×:Ò:¸;ÑGÔG‰ˆ�aØ# GÔ,ˆð      4 Ô(ˆ
Ø# A A A t GÔ,ˆõ Œ]‰_Œ_ð 		ið 		iØ $× DÒ DØØAÐAÐAÐAØ”ØÔ%ØÔ,ñ!ô !Ðõ (¨Ð6GÐW\Ð]Ñ]Ô]×eÒeÐfgÑhÔhˆLð		ið 		ið 		iñ 		iô 		ið 		ið 		ið 		ið 		ið 		ið 		iøøøð 		ið 		ið 		ið 		iõ $ JÐ0AÐQVÐWÑWÔW×_Ò_Ð`aÑbÔbˆõ 4°LÀ,ÐPYÑZÔZÝ" <°¸yÑIÔIð
ð 
ˆð
 ØØˆs   Á?ACÃC Ã#C c                 óœ   — t          j        d„ t          |¦  «        D ¦   «         ¦  «        }t          j        d„ |D ¦   «         ¦  «        }||fS )Nc                 óD   — g | ]\  }\  }}t          j        ||¦  «        ‘ŒS r
  ©rF   Ú	full_like)r<  r=  Úsrcrø   s       r1   r?  zDVideomtLoss._get_predictions_permutation_indices.<locals>.<listcomp>  s,   € Ð"aÐ"aÐ"a¹{¸qÁ(À3È¥5¤?°3¸Ñ#:Ô#:Ð"aÐ"aÐ"ar2   c                 ó   — g | ]\  }}|‘ŒS r
  r
  )r<  r¦  rø   s      r1   r?  zDVideomtLoss._get_predictions_permutation_indices.<locals>.<listcomp>  s   € Ð(EÐ(EÐ(E±°#°q¨Ð(EÐ(EÐ(Er2   ©rF   rl   rm  )r.   rH  Úbatch_indicesÚpredictions_indicess       r1   r‡  z0VideomtLoss._get_predictions_permutation_indices
  sT   € åœ	Ð"aÐ"aÍiÐX_ÑN`ÔN`Ð"aÑ"aÔ"aÑbÔbˆÝ#œiÐ(EÐ(E¸WÐ(EÑ(EÔ(EÑFÔFÐØÐ1Ð1Ð1r2   c                 óœ   — t          j        d„ t          |¦  «        D ¦   «         ¦  «        }t          j        d„ |D ¦   «         ¦  «        }||fS )Nc                 óD   — g | ]\  }\  }}t          j        ||¦  «        ‘ŒS r
  r¤  )r<  r=  rø   Útgts       r1   r?  z@VideomtLoss._get_targets_permutation_indices.<locals>.<listcomp>  s,   € Ð"aÐ"aÐ"a¹{¸qÁ(À1Àc¥5¤?°3¸Ñ#:Ô#:Ð"aÐ"aÐ"ar2   c                 ó   — g | ]\  }}|‘ŒS r
  r
  )r<  rø   r­  s      r1   r?  z@VideomtLoss._get_targets_permutation_indices.<locals>.<listcomp>  s   € Ð#@Ð#@Ð#@©H¨Q° CÐ#@Ð#@Ð#@r2   r¨  )r.   rH  r©  Útarget_indicess       r1   r™  z,VideomtLoss._get_targets_permutation_indices  sR   € åœ	Ð"aÐ"aÍiÐX_ÑN`ÔN`Ð"aÑ"aÔ"aÑbÔbˆÝœÐ#@Ð#@¸Ð#@Ñ#@Ô#@ÑAÔAˆØ˜nÐ,Ð,r2   r•  c                 ó0   — t          j        |¦  «         }|S )a…  
        In Videomt paper, uncertainty is estimated as L1 distance between 0.0 and the logit prediction in 'logits'
        for the foreground class in `classes`.

        Args:
            logits (`torch.Tensor`):
            A tensor of shape (R, 1, ...) for class-specific or class-agnostic, where R is the total number of predicted masks in all images and C is:
            the number of foreground classes. The values are logits.

        Returns:
            scores (`torch.Tensor`): A tensor of shape (R, 1, ...) that contains uncertainty scores with the most
            uncertain locations having the highest uncertainty score.
        )rF   Úabs)r.   r•  Úuncertainty_scoress      r1   r”  z!VideomtLoss.calculate_uncertainty  s   € õ  %œy¨Ñ0Ô0Ð1ÐØ!Ð!r2   r1  re  rf  c           	      ó¬  — |j         d         }t          ||z  ¦  «        }t          j        ||d|j        ¬¦  «        }t          ||d¬¦  «        }	 ||	¦  «        }
t          ||z  ¦  «        }||z
  }t          j        |
dd…ddd…f         |d¬¦  «        d         }|t          j        |t          j        |j        ¬	¦  «        z  }||dd…df         z  }| 	                    d
d¦  «        | 	                    d
¦  «        dd…f          	                    ||d¦  «        }|dk    r3t          j
        |t          j        ||d|j        ¬¦  «        gd¬¦  «        }|S )a€  
        This function is meant for sampling points in [0, 1] * [0, 1] coordinate space based on their uncertainty. The
        uncertainty is calculated for each point using the passed `uncertainty function` that takes points logit
        prediction as input.

        Args:
            logits (`float`):
                Logit predictions for P points.
            uncertainty_function:
                A function that takes logit predictions for P points and returns their uncertainties.
            num_points (`int`):
                The number of points P to sample.
            oversample_ratio (`int`):
                Oversampling parameter.
            importance_sample_ratio (`float`):
                Ratio of points that are sampled via importance sampling.

        Returns:
            point_coordinates (`torch.Tensor`):
                Coordinates for P sampled points.
        r   r9   r6  Fr7  Nr   )Úkrf   rØ   rN   re   )r:   rz   rF   rÙ   rd   r  Útopkr^   ÚlongrÅ   rl   )r.   r•  Úuncertainty_functionr1  re  rf  Ú	num_boxesÚnum_points_sampledr  r   Úpoint_uncertaintiesÚnum_uncertain_pointsÚnum_random_pointsr‹  Úshifts                  r1   rš  z+VideomtLoss.sample_points_using_uncertainty'  s�  € ð< ”L ”Oˆ	Ý  Ð.>Ñ!>Ñ?Ô?Ðõ "œJ yÐ2DÀaÐPVÔP]Ð^Ñ^Ô^Ðå# FÐ,=ÈUÐSÑSÔSˆà2Ð2°<Ñ@Ô@Ðå"Ð#:¸ZÑ#GÑHÔHÐØ&Ð)=Ñ=ÐåŒjÐ,¨Q¨Q¨Q°°1°1°1¨WÔ5Ð9MÐSTÐUÑUÔUÐVWÔXˆØ"¥U¤\°)Å5Ä:ÐV\ÔVcÐ%dÑ%dÔ%dÑdˆØˆu�Q�Q�Q˜�WŒ~ÑˆØ-×2Ò2°2°qÑ9Ô9¸#¿(º(À2¹,¼,ÈÈÈ¸/ÔJ×OÒOÐPYÐ[oÐqrÑsÔsÐà˜qÒ Ð Ý %¤	Ø"¥E¤J¨yÐ:KÈQÐW]ÔWdÐ$eÑ$eÔ$eÐfØð!ñ !ô !Ðð !Ð r2   NÚauxiliary_predictionsc                 óÆ  ‡— |                       ||||¦  «        }|                      ||d         j        ¬¦  «        }i |                      ||||¦  «        ¥|                      |||¦  «        ¥}|�rt          |¦  «        D ]b\  Š}	|	d         }|	d         }|                      ||||¦  «        }
ˆfd„|
                     ¦   «         D ¦   «         }
|                     |
¦  «         Œc|S )a»  
        This performs the loss computation.

        Args:
            masks_queries_logits (`torch.Tensor`):
                A tensor of shape `(batch_size, num_queries, height, width)`.
            class_queries_logits (`torch.Tensor`):
                A tensor of shape `(batch_size, num_queries, num_labels)`.
            mask_labels (`torch.Tensor`):
                List of mask labels of shape `(labels, height, width)`.
            class_labels (`list[torch.Tensor]`):
                List of class labels of shape `(labels)`.
            auxiliary_predictions (`dict[str, torch.Tensor]`, *optional*):
                if `use_auxiliary_loss` was set to `true` in [`VideomtConfig`], then it contains the logits from
                the inner layers of the VideomtMaskedAttentionDecoder.

        Returns:
            losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing three keys:
            - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels.
            - **loss_mask** -- The loss computed using sigmoid cross_entropy loss on the predicted and ground truth
              masks.
            - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth
              masks.
            if `use_auxiliary_loss` was set to `true` in [`VideomtConfig`], the dictionary contains additional
            losses for each auxiliary predictions.
        r   r6  Nr  r  c                 ó&   •— i | ]\  }}|› d ‰› �|“ŒS )rø   r
  )r<  r¤   r¥   r‹  s      €r1   ú
<dictcomp>z'VideomtLoss.forward.<locals>.<dictcomp>�  s)   ø€ ÐWÐWÐW±z°s¸E ˜^˜^ c˜^˜^¨UÐWÐWÐWr2   )	rj  Úget_num_masksrd   r¡  r‘  rm  rA   ÚitemsÚupdate)r.   r  r  r3  r4  r¾  rH  rS  r�  Úaux_outputsÚ	loss_dictr‹  s              @r1   rA   zVideomtLoss.forward^  s  ø€ ðH —,’,Ð3Ð5IÈ;ÐXdÑeÔeˆà×&Ò& |¸LÈ¼OÔ<RÐ&ÑSÔSˆ	ð%
Ø�oŠoÐ2°KÀÈ)ÑTÔTð%
à×ÒÐ3°\À7ÑKÔKð%
ˆð
 !Ð,Ý$-Ð.CÑ$DÔ$Dð )ð )Ñ ��[Ø'2Ð3IÔ'JÐ$Ø'2Ð3IÔ'JÐ$Ø ŸLšLÐ)=Ð?SÐU`ÐbnÑoÔo�	ØWÐWÐWÐWÀYÇ_Â_ÑEVÔEVÐWÑWÔW�	Ø—’˜iÑ(Ô(Ð(Ð(àˆr2   rd   c                 ó0  — t          d„ |D ¦   «         ¦  «        }t          j        |t          j        |¬¦  «        }d}t	          ¦   «         r2t
          j        i k    r"t          |¦  «        }t          ¦   «         j        }t          j	        ||z  d¬¦  «        }|S )zk
        Computes the average number of target masks across the batch, for normalization purposes.
        c              3   ó4   K  — | ]}t          |¦  «        V — Œd S rs   )rw  )r<  Úclassess     r1   ú	<genexpr>z,VideomtLoss.get_num_masks.<locals>.<genexpr>™  s(   è è € ÐAÐA¨�˜G™œÐAÐAÐAÐAÐAÐAr2   rØ   r   )Úmin)
r  rF   r:  râ   r   r   Ú_shared_stater   Únum_processesÚclamp)r.   r4  rd   rS  Ú
world_sizes        r1   rÂ  zVideomtLoss.get_num_masks•  s‘   € õ ÐAÐA°LÐAÑAÔAÑAÔAˆ	Ý”O IµU´[ÈÐPÑPÔPˆ	Øˆ
Ý"Ñ$Ô$ð 	:ÝÔ)¨RÒ/Ð/Ý" 9Ñ-Ô-�	Ý)™^œ^Ô9�
å”K 	¨JÑ 6¸AÐ>Ñ>Ô>ˆ	ØÐr2   rs   )rB   rC   rD   r   Údictr   râ   r"   rR  rz   rs  r   rÎ   r€  ÚnpÚarrayr‘  rF   r¡  r‡  r™  r”  rš  rA   rd   rÂ  rG   rH   s   @r1   r\  r\  m  s¦  ø€ € € € € ð!
˜}ð !
¸4ÀÀUÀ
Ô;Kð !
ð !
ð !
ð !
ð !
ð !
ðF $ t¨C¤y¤/ð °d¸3´ið ð ð ð ð-°4¸´<ð -ÀEÈ&ÐRXÈ.ÔDYð -ð -ð -ð -ð" Ø$*ð Ø:>¸v¼,ð ØQVÐWYÔW_ÔQ`ð à	ˆc�6ˆkÔ	ð ð  ð  ð  ðD<à#œlð<ð ˜%œ,Ô'ð<ð �r”x”ð	<ð
 ð<ð 
ˆc�5”<ÐÔ	 ð<ð <ð <ð <ð|2ð 2ð 2ð-ð -ð -ð"¨E¬Lð "¸U¼\ð "ð "ð "ð "ð"5!à”ð5!ð ð	5!ð
 ð5!ð "'ð5!ð 
Œð5!ð 5!ð 5!ð 5!ðz AEð5ð 5à#œlð5ð $œlð5ð ˜%œ,Ô'ð	5ð
 ˜5œ<Ô(ð5ð  $ C¨¬Ð$5Ô6¸Ñ=ð5ð 
ˆc�5”<ÐÔ	 ð5ð 5ð 5ð 5ðn¨%¬,ð ÀÄð ÐQVÔQ]ð ð ð ð ð ð ð ð r2   r\  c                   ó�   ‡ — e Zd ZU dZeed<   dZdZdZdZ	dgZ
dZeed	œZ ej        ¦   «         d
ej        ddfˆ fd„¦   «         Zˆ xZS )ÚVideomtPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    r/   ÚvideomtÚpixel_values_videos)ÚvideoFrä   T)rÃ   r  r¢   r4   Nc                 ó8  •— t          ¦   «                              |¦  «         | j        j        }t	          |t
          j        t
          j        t
          j        f¦  «        r�t          j
        |j        t          j        d¦  «        ¬¦  «         |j        �gt          j        j	                             |j        ¦  «        \  }}|dk    rdt          j        |¦  «        z  nd}t          j        |j        | |¦  «         �nât	          |t
          j        ¦  «        r_t          j        |j        dd¬¦  «         |j        �:t+          |j        dd¦  «        s$t          j        |j        |j                 ¦  «         �nit	          |t.          ¦  «        r6t1          |d	¦  «        r$t          j        |j        | j        j        ¦  «         �nt	          |t8          ¦  «        r…t          j        |j        d|¬¦  «         t          j        |j        ¦  «         t          j         |j!        t          j"        |j!        j#        d
         ¦  «         $                    d¦  «        ¦  «         n„t	          |tJ          ¦  «        rAt          j&        |j'        dz   ¦  «        }|j(        |d
<   t          j         |j)        |¦  «         n.t	          |tT          ¦  «        rt          j+        |j,        ¦  «         t	          |t8          ¦  «        r&t
          j	                             |j-        ¦  «         d S d S )Nrc   )Úar   r   r¡   )rX  ÚstdÚ_is_hf_initializedFrþ   rN   rM   ).r!   Ú_init_weightsr/   Úinitializer_ranger'   r   r|   r,   ÚConvTranspose2dÚinitÚkaiming_uniform_r=   ÚmathÚsqrtry   rF   Ú_calculate_fan_in_and_fan_outÚuniform_r[   Únormal_Úpadding_idxÚgetattrÚzeros_rì   ÚhasattrÚ	constant_rþ   rü   rJ   Útrunc_normal_rR   rU   ry  rL   r^   r:   r_   r\  rý   ra  rc  r`  ÚVideomtForUniversalSegmentationÚones_Úattn_mask_probsr`   )r.   r¢   rÚ  Úfan_inrø   Úboundr`  r0   s          €r1   rÜ  z$VideomtPreTrainedModel._init_weights¸  s•  ø€ å‰Œ×Ò˜fÑ%Ô%Ð%ØŒkÔ+ˆÝ�f�rœy­"¬)µRÔ5GÐHÑIÔIð 	/ÝÔ! &¤-µ4´9¸Q±<´<Ð@Ñ@Ô@Ð@ØŒ{Ð&Ý!œHœM×GÒGÈÌÑVÔV‘	�˜Ø17¸!²°˜�DœI fÑ-Ô-Ñ-Ð-À�Ý”˜fœk¨E¨6°5Ñ9Ô9Ð9ùÝ˜¥¤Ñ-Ô-ð 	/ÝŒL˜œ¨S°aÐ8Ñ8Ô8Ð8àÔ!Ð-µg¸f¼mÐMaÐchÑ6iÔ6iÐ-Ý”˜FœM¨&Ô*<Ô=Ñ>Ô>Ð>ùÝ˜Õ 1Ñ2Ô2ð 	/Ý�v˜yÑ)Ô)ð MÝ”˜vœ~¨t¬{Ô/KÑLÔLÐLùÝ˜Õ 1Ñ2Ô2ð 		/ÝÔ˜vÔ/°c¸sÐCÑCÔCÐCÝŒK˜Ô.Ñ/Ô/Ð/ÝŒJ�vÔ*­E¬L¸Ô9LÔ9RÐSUÔ9VÑ,WÔ,W×,^Ò,^Ð_fÑ,gÔ,gÑhÔhÐhÐhÝ˜¥Ñ,Ô,ð 	/Ý œ: fÔ&7¸!Ñ&;Ñ<Ô<ˆLØ%œˆL˜ÑÝŒJ�vÔ*¨LÑ9Ô9Ð9Ð9Ý˜Õ ?Ñ@Ô@ð 	/ÝŒJ�vÔ-Ñ.Ô.Ð.Ý�fÕ/Ñ0Ô0ð 	.ÝŒG�NŠN˜6Ô,Ñ-Ô-Ð-Ð-Ð-ð	.ð 	.r2   )rB   rC   rD   rE   r   r	  Úbase_model_prefixÚmain_input_nameÚinput_modalitiesÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_supports_sdparä   rµ   Ú_can_record_outputsrF   rQ  r   ÚModulerÜ  rG   rH   s   @r1   rÔ  rÔ  ¥  s³   ø€ € € € € € ðð ð
 ÐÐÑØ!ÐØ+€OØ!ÐØ&+Ð#Ø'Ð(ÐØ€Nà%Ø&ðð Ðð
 €U„]�_„_ð. B¤Ið .°$ð .ð .ð .ð .ð .ñ „_ð.ð .ð .ð .ð .r2   rÔ  c                   óD   ‡ — e Zd Zdˆ fd„	Zdej        dej        fd„Zˆ xZS )ÚVideomtLayerNorm2dç�íµ ÷Æ°>Tc                 óP   •— t          ¦   «                              |||¬¦  «         d S )N)rç   Úelementwise_affine)r!   r"   )r.   r%   rç   Úaffiner0   s       €r1   r"   zVideomtLayerNorm2d.__init__Ù  s(   ø€ Ý‰Œ×Ò˜¨3À6ÐÑJÔJÐJÐJÐJr2   r†   r4   c                 ó¾   — |                      dddd¦  «        }t          j        || j        | j        | j        | j        ¦  «        }|                      dddd¦  «        }|S )Nr   r9   r   r   )ÚpermuteÚFÚ
layer_normÚnormalized_shaper=   ry   rç   rˆ   s     r1   rA   zVideomtLayerNorm2d.forwardÜ  s^   € Ø#×+Ò+¨A¨q°!°QÑ7Ô7ˆÝ”| L°$Ô2GÈÌÐVZÔV_ÐaeÔaiÑjÔjˆØ#×+Ò+¨A¨q°!°QÑ7Ô7ˆØÐr2   )rû  TrŠ   rH   s   @r1   rú  rú  Ø  si   ø€ € € € € ðKð Kð Kð Kð Kð Kð E¤Lð °U´\ð ð ð ð ð ð ð ð r2   rú  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚVideomtScaleLayerr/   c                 ó$  •— t          ¦   «                              ¦   «          |j        }t          j        ||dd¬¦  «        | _        t          |j                 | _        t          j	        ||dd|d¬¦  «        | _
        t          |¦  «        | _        d S )Nr9   r   r   r   F)r   ÚpaddingÚgroupsry   )r!   r"   r&   r   rÞ  Úconv1r	   r~   r€   r,   Úconv2rú  Úlayernorm2d©r.   r/   r&   r0   s      €r1   r"   zVideomtScaleLayer.__init__ä  sŽ   ø€ Ý‰Œ×ÒÑÔÐØÔ(ˆÝÔ'¨°[ÈaÐXYÐZÑZÔZˆŒ
Ý  Ô!2Ô3ˆŒÝ”YØØØØØØð
ñ 
ô 
ˆŒ
õ .¨kÑ:Ô:ˆÔÐÐr2   rÃ   r4   c                 ó®   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|S rs   )r	  r€   r
  r  ©r.   rÃ   s     r1   rA   zVideomtScaleLayer.forwardô  sN   € ØŸ
š
 =Ñ1Ô1ˆØŸš¨Ñ6Ô6ˆØŸ
š
 =Ñ1Ô1ˆØ×(Ò(¨Ñ7Ô7ˆØÐr2   ©	rB   rC   rD   r   r"   rF   r   rA   rG   rH   s   @r1   r  r  ã  sj   ø€ € € € € ð;˜}ð ;ð ;ð ;ð ;ð ;ð ;ð  U¤\ð °e´lð ð ð ð ð ð ð ð r2   r  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚVideomtScaleBlockr/   c                 óÐ   •‡— t          ¦   «                              ¦   «          ‰j        | _        t	          j        ˆfd„t          | j        ¦  «        D ¦   «         ¦  «        | _        d S )Nc                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r
  )r  ©r<  rø   r/   s     €r1   r?  z.VideomtScaleBlock.__init__.<locals>.<listcomp>   s"   ø€ Ð#^Ð#^Ð#^À!Õ$5°fÑ$=Ô$=Ð#^Ð#^Ð#^r2   )r!   r"   Únum_upscale_blocksÚ
num_blocksr   Ú
ModuleListr@  ÚblockrÂ   s    `€r1   r"   zVideomtScaleBlock.__init__ý  sX   øø€ Ý‰Œ×ÒÑÔÐØ Ô3ˆŒÝ”]Ð#^Ð#^Ð#^Ð#^ÅuÈTÌ_ÑG]ÔG]Ð#^Ñ#^Ô#^Ñ_Ô_ˆŒ
ˆ
ˆ
r2   rÃ   r4   c                 ó0   — | j         D ]} ||¦  «        }Œ|S rs   )r  )r.   rÃ   r  s      r1   rA   zVideomtScaleBlock.forward  s*   € Ø”Zð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMØÐr2   r  rH   s   @r1   r  r  ü  sq   ø€ € € € € ð`˜}ð `ð `ð `ð `ð `ð `ð
 U¤\ð °e´lð ð ð ð ð ð ð ð r2   r  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚVideomtMaskHeadr/   c                 ó   •— t          ¦   «                              ¦   «          |j        }t          j        ||¦  «        | _        t          j        ||¦  «        | _        t          j        ||¦  «        | _        t          |j	                 | _
        d S rs   )r!   r"   r&   r   r|   r}   r�   Úfc3r	   r~   r€   r  s      €r1   r"   zVideomtMaskHead.__init__	  sm   ø€ Ý‰Œ×ÒÑÔÐàÔ(ˆÝ”9˜[¨+Ñ6Ô6ˆŒÝ”9˜[¨+Ñ6Ô6ˆŒÝ”9˜[¨+Ñ6Ô6ˆŒÝ  Ô!2Ô3ˆŒˆˆr2   rÃ   r4   c                 óÐ   — |                       |                      |¦  «        ¦  «        }|                       |                      |¦  «        ¦  «        }|                      |¦  «        }|S rs   )r€   r}   r�   r  r  s     r1   rA   zVideomtMaskHead.forward  sS   € ØŸš¨¯ª°Ñ(?Ô(?Ñ@Ô@ˆØŸš¨¯ª°Ñ(?Ô(?Ñ@Ô@ˆØŸš Ñ/Ô/ˆØÐr2   r  rH   s   @r1   r  r    sj   ø€ € € € € ð4˜}ð 4ð 4ð 4ð 4ð 4ð 4ð U¤\ð °e´lð ð ð ð ð ð ð ð r2   r  zY
    The Videomt Model with head on top for instance/semantic/panoptic segmentation.
    c                   ón  ‡ — e Zd ZdZdefˆ fd„Zdededededeeef         d	eeef         fd
„Z	deeef         d	efd„Z
eee	 	 	 	 ddej        dz  deej                 dz  deej                 dz  deej                 dz  dee         d	efd„¦   «         ¦   «         ¦   «         Zd„ Zdej        fd„Zˆ xZS )rì  rÖ  r/   c                 ó²  •‡— t          ¦   «                              ‰¦  «         ‰| _        ‰j        | _        t	          ‰¦  «        | _        t          j        ‰j        ‰j	        ¬¦  «        | _
        t          j        ‰j        ‰j        ¦  «        | _        t          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t#          ‰¦  «        | _        t'          ‰¦  «        | _        t          j        ‰j        ‰j        dz   ¦  «        | _        ‰j        ‰j        z  ‰j        ‰j        z  f| _        ‰j        ‰j        ‰j        dœ| _        t?          ‰| j        ¬¦  «        | _         |  !                    dtE          j#        ‰j$        ¦  «        ¦  «         t          j        ‰j        ‰j        ¦  «        | _%        |  &                    ¦   «          d S )Nræ   c                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r
  )rä   r  s     €r1   r?  z<VideomtForUniversalSegmentation.__init__.<locals>.<listcomp>)  s!   ø€ Ð$cÐ$cÐ$c¸a¥\°&Ñ%9Ô%9Ð$cÐ$cÐ$cr2   r   )r…  r—  r˜  )r/   r]  rî  )'r!   r"   r/   Únum_hidden_layersrJ   r@   r   rè   r&   ré   Ú	layernormr[   rŠ  r£   r  r@  Úlayersr  Úupscale_blockr  Ú	mask_headr|   ra  Úclass_predictorr#   r$   Ú	grid_sizerg  ri  rh  r]  r\  r%  r]   rF   rý   r  Úquery_updaterÚ	post_initrÂ   s    `€r1   r"   z(VideomtForUniversalSegmentation.__init__!  s“  øø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒØ!'Ô!9ˆÔÝ+¨FÑ3Ô3ˆŒÝœ fÔ&8¸fÔ>SÐTÑTÔTˆŒå”\ &Ô"4°fÔ6HÑIÔIˆŒ
Ý”mÐ$cÐ$cÐ$cÐ$cÅ5ÈÔIaÑCbÔCbÐ$cÑ$cÔ$cÑdÔdˆŒå.¨vÑ6Ô6ˆÔÝ(¨Ñ0Ô0ˆŒå!œy¨Ô);¸VÔ=NÐQRÑ=RÑSÔSˆÔà Ô+¨vÔ/@Ñ@À&ÔBSÐW]ÔWhÑBhÐiˆŒà"(Ô"5ØÔ+ØÔ+ð.
ð .
ˆÔõ %¨FÀÔ@PÐQÑQÔQˆŒà×ÒÐ.µ´
¸6Ô;LÑ0MÔ0MÑNÔNÐNÝœY vÔ'9¸6Ô;MÑNÔNˆÔà�ŠÑÔÐÐÐr2   r  r  r3  r4  r¾  r4   c                 ó¾   — |                       |||||¬¦  «        }| j                             ¦   «         D ](\  }}|                     ¦   «         D ]\  }	}
||	v r|
|z  }
ŒŒ)|S )N)r  r  r3  r4  r¾  )r%  r]  rÃ  )r.   r  r  r3  r4  r¾  rÆ  r¤   r=   Úloss_keyr  s              r1   Úget_loss_dictz-VideomtForUniversalSegmentation.get_loss_dict>  sŒ   € ð (,§~¢~Ø!5Ø!5Ø#Ø%Ø"7ð (6ñ (
ô (
ˆ	ð  Ô+×1Ò1Ñ3Ô3ð 	#ð 	#‰KˆC�Ø"+§/¢/Ñ"3Ô"3ð #ð #‘�˜$Ø˜(�?�?Ø˜F‘N�Døð#ð Ðr2   rÆ  c                 óD   — t          |                     ¦   «         ¦  «        S rs   )r  rÌ   )r.   rÆ  s     r1   Úget_lossz(VideomtForUniversalSegmentation.get_lossV  s   € Ý�9×#Ò#Ñ%Ô%Ñ&Ô&Ð&r2   NÚpatch_offsetsr°   c           	      ó’  — d|v rt          d¦  «        ‚|€t          d¦  «        ‚|j        dk    rt          d¦  «        ‚|€|�t          d¦  «        ‚|j        \  }}}}	}
|                     ||z  ||	|
¦  «        }|                      |¦  «        }| j        | j        j        z
  }| j        d|…         D ]} ||¦  «        }Œ| 	                    |||j        d         |j        d	         ¦  «        }g }g }g }d}t          |¦  «        D �]s}|dd…|f         }|€G| j        j        ddd…dd…f                              |d
d
¦  «                             |j        ¦  «        }n_|                      |¦  «                             |j        ¦  «        | j        j        ddd…dd…f                              |j        ¦  «        z   }t#          j        ||fd¬¦  «        }| j        |d…         D ]} ||¦  «        }Œ|                      |¦  «        }|                      |¦  «        \  }}|                     |¦  «         |                     |¦  «         |                     |¦  «         |dd…d| j        j        …dd…f         }�Œut/          dt#          j        |d¬¦  «        t#          j        |d¬¦  «        t#          j        |d¬¦  «        ¬¦  «        S )aø  
        pixel_values_videos (`torch.Tensor`, *optional*):
            Video inputs of shape `(batch_size, num_frames, num_channels, height, width)`.
        mask_labels (`list[torch.Tensor]`, *optional*):
            Not supported for 5D video inputs.
        class_labels (`list[torch.LongTensor]`, *optional*):
            Not supported for 5D video inputs.
        patch_offsets (`list[torch.Tensor]`, *optional*):
            Unused for video inputs and only kept for modular compatibility.
        r3   zAUse `pixel_values_videos` with `VideomtForUniversalSegmentation`.Nz'You have to specify pixel_values_videosrc   zyVideomtForUniversalSegmentation only supports 5D video inputs of shape (batch_size, num_frames, channels, height, width).z“Training with 5D video inputs is not supported in `VideomtForUniversalSegmentation`. Flatten frames and use `EomtForUniversalSegmentation` instead.r   r9   rN   re   r   )r  r  r  r  )r;   rg   r:   rh   r@   r"  r/   r  r$  rÅ   r@  r£   r=   r_   r<   rd   r)  rF   rl   r#  ÚpredictrG  rŠ  r  )r.   rÖ  r3  r4  r0  r°   rm   rn   r%   ro   rp   Úflat_pixel_valuesrÃ   Úquery_start_idxÚlayer_moduleÚall_masks_queries_logitsÚall_class_queries_logitsÚall_last_hidden_statesÚpropagated_queryÚ	frame_idxÚframe_hidden_statesÚquery_tokensÚsequence_outputr  r  s                            r1   rA   z'VideomtForUniversalSegmentation.forwardY  sT  € ð* ˜VÐ#Ð#ÝÐ`ÑaÔaÐaàÐ&ÝÐFÑGÔGÐGàÔ# qÒ(Ð(ÝðEñô ð ð
 Ð" lÐ&>ÝðQñô ð ð
 ?RÔ>WÑ;ˆ
�J ¨f°eØ/×7Ò7¸
ÀZÑ8OÐQ]Ð_eÐglÑmÔmÐàŸšÐ(9Ñ:Ô:ˆØÔ0°4´;Ô3IÑIˆà œKÐ(8¨Ð(8Ô9ð 	8ð 	8ˆLØ(˜L¨Ñ7Ô7ˆMˆMà%×*Ò*¨:°zÀ=ÔCVÐWXÔCYÐ[hÔ[nÐopÔ[qÑrÔrˆà#%Ð Ø#%Ð Ø!#ÐØÐå˜zÑ*Ô*ð 	Tñ 	TˆIØ"/°°°°9°Ô"=ÐàÐ'Ø#œzÔ0°°q°q°q¸!¸!¸!°Ô<×CÒCÀJÐPRÐTVÑWÔW×ZÒZÐ[nÔ[uÑvÔv��à#×1Ò1Ð2BÑCÔC×FÒFÐGZÔGaÑbÔbÐeiÔeoÔevØ˜!˜!˜!˜Q˜Q˜Q�Jôfç’"Ð(Ô/Ñ0Ô0ñ 1�õ #(¤)¨\Ð;NÐ,OÐUVÐ"WÑ"WÔ"WÐà $¤¨OÐ,<Ð,<Ô =ð Hð H�Ø&2 lÐ3FÑ&GÔ&GÐ#Ð#à"ŸnšnÐ-@ÑAÔAˆOØ9=¿ºÀoÑ9VÔ9VÑ6Ð Ð"6à$×+Ò+Ð,@ÑAÔAÐAØ$×+Ò+Ð,@ÑAÔAÐAØ"×)Ò)¨/Ñ:Ô:Ð:Ø2°1°1°1Ð6O¸¼Ô8OÐ6OÐQRÐQRÐQRÐ3RÔSÐÑå4ØÝ!&¤Ð+CÈÐ!KÑ!KÔ!KÝ!&¤Ð+CÈÐ!KÑ!KÔ!KÝ#œiÐ(>ÀAÐFÑFÔFð	
ñ 
ô 
ð 	
r2   c                 ó   — | j         j        S rs   )r@   rV   rß   s    r1   Úget_input_embeddingsz4VideomtForUniversalSegmentation.get_input_embeddings­  s   € ØŒÔ/Ð/r2   r•  c                 ó¤  — |d d …d | j         j        …d d …f         }|                      |¦  «        }|d d …| j         j        | j        j        z   d …d d …f         }|                     dd¦  «        } |j        |j        d         dg| j        ¢R Ž }|  	                    |¦  «        }|  
                    |¦  «        }t          j        d||¦  «        }||fS )Nr   r9   r   rN   zbqc, bchw -> bqhw)r/   rŠ  r'  r@   rZ   r?   rh   r:   r(  r&  r%  rF   Úeinsum)r.   r•  r<  Úclass_logitsÚprefix_tokensÚmask_logitss         r1   r2  z'VideomtForUniversalSegmentation.predict°  sè   € Ø˜a˜a˜aÐ!: 4¤;Ô#:Ð!:¸A¸A¸AÐ=Ô>ˆØ×+Ò+¨LÑ9Ô9ˆà˜q˜q˜q $¤+Ô"9¸D¼OÔ<]Ñ"]Ð"_Ð"_ÐabÐabÐabÐbÔcˆØ%×/Ò/°°1Ñ5Ô5ˆà-˜Ô-¨mÔ.AÀ!Ô.DÀbÐZÈ4Ì>ÐZÐZÐZˆà—~’~ lÑ3Ô3ˆØ×*Ò*¨=Ñ9Ô9ˆå”lÐ#6¸ÀmÑTÔTˆà˜LÐ(Ð(r2   )NNNN)rB   rC   rD   rò  r   r"   r   rÐ  r   r-  r/  r   r   r   rF   rR  r   r   r  rA   r?  r2  rG   rH   s   @r1   rì  rì    sÁ  ø€ € € € € ð ,€Oð˜}ð ð ð ð ð ð ð:à$ðð %ðð ð	ð
 ðð  $ C¨ KÔ0ðð 
ˆc�6ˆkÔ	ðð ð ð ð0' $ s¨F {Ô"3ð '¸ð 'ð 'ð 'ð 'ð  ØØð 48Ø15Ø26Ø37ðO
ð O
à"œ\¨DÑ0ðO
ð ˜%œ,Ô'¨$Ñ.ðO
ð ˜5œ<Ô(¨4Ñ/ð	O
ð
 ˜EœLÔ)¨DÑ0ðO
ð Ð+Ô,ðO
ð 
/ðO
ð O
ð O
ñ „^ñ „_ñ  ÔðO
ðb0ð 0ð 0ð)˜eœlð )ð )ð )ð )ð )ð )ð )ð )r2   rì  rá   )F)KÚcollections.abcr(   rá  r   Údataclassesr   ÚnumpyrÑ  rF   Útorch.nn.functionalr   r™   r  r   Ú r   rß  Úactivationsr	   Ú
file_utilsr
   r   r   Úmodeling_layersr   Úmodeling_utilsr   r   Úprocessing_utilsr   Úutilsr   r   r   Úutils.genericr   Úutils.output_capturingr   Úconfiguration_videomtr   Úscipy.optimizer   Ú
accelerater   Úaccelerate.utilsr   rø  r   rJ   rv   rŒ   râ   r³   rµ   rÐ   rÔ   rä   rì   r  r  r  r*  r,  rz   rV  rZ  r\  rÔ  rè   rú  r  r  r  rì  Ú__all__r
  r2   r1   ú<module>rW     sÊ  ðð* Ð Ð Ð Ø €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LØ 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø PÐ PÐ PÐ PÐ PÐ PÐ PÐ PÐ PÐ PØ 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø 0Ð 0Ð 0Ð 0Ð 0Ð 0ð ÐÑÔð 5Ø4Ð4Ð4Ð4Ð4Ð4àÐÑÔð (Ø'Ð'Ð'Ð'Ð'Ð'Ø'Ð'Ð'Ð'Ð'Ð'ð ð  ð  ð  ð  ˜RœYñ  ô  ð  ðF,ð ,ð ,ð ,ð ,˜œ	ñ ,ô ,ð ,ð^ð ð ð ð �”ñ ô ð ð&(ð (ð (ð (ð (�b”iñ (ô (ð (ð0 ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð ð%ð ð%ð %ð %ð %ð.8)ð 8)ð 8)ð 8)ð 8)�r”yñ 8)ô 8)ð 8)ðv(ð (ð (ð (ð (�r”yñ (ô (ð (ð"%ð %ð %ð %ð %�b”iñ %ô %ð %ð0'ð 'ð 'ð 'ð 'Ð-ñ 'ô 'ð 'ðT+ð +ð +ð +ð +˜œ	ñ +ô +ð +ð €ðð	ñ 	ô 	ð ð7ð 7ð 7ð 7ð 7¨Kñ 7ô 7ñ „ñ	ô 	ð7ð< LQðð Ø”LðØ5:´\ðà
„\ðð ð ð ð@ ð °ð ¸6ð ð ð ð ð,°´ð ÀuÄ|ð ÐX]ÔXdð ð ð ð ð8gð gð gð gð g˜bœiñ gô gð gðT�fð  fð ¸ð Àð ð ð ð ð< u¤|ð ¸U¼\ð ÐVYð Ð^cÔ^jð ð ð ð ð(uð uð uð uð u�"”)ñ uô uð uðp	 ð/.ð /.ð /.ð /.ð /.˜_ñ /.ô /.ñ „ð/.ðdð ð ð ð ˜œñ ô ð ðð ð ð ð ˜œ	ñ ô ð ð2	ð 	ð 	ð 	ð 	˜œ	ñ 	ô 	ð 	ðð ð ð ð �b”iñ ô ð ð" €ððñ ô ð
`)ð `)ð `)ð `)ð `)Ð&<ñ `)ô `)ñô ð
`)ðF $Ð%FÐ
G€€€r2   