§
    ‚Štj‰î  ã                   ó²  — d dl Z d dlmZ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 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 m!Z!m"Z" ddl#m$Z$ ddl%m&Z& ddl'm(Z(m)Z) ddl*m+Z+m,Z,m-Z- ddl.m/Z/ ddl0m1Z1  e)d¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z2 e)d¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z3 e)d¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z4 e)d¬¦  «        e G d„ d e¦  «        ¦   «         ¦   «         Z5 e)d!¬¦  «        e G d"„ d#e¦  «        ¦   «         ¦   «         Z6 G d$„ d%e
j7        ¦  «        Z8 e&d&¬'¦  «        d(e9d)e9d*ej:        d+ej;        d,ej        f
d-„¦   «         Z<	 	 	 d�d.ej        d/e=dz  d0e=dz  d1e=dz  d,ej        f
d2„Z> G d3„ d4e
j7        ¦  «        Z? ed5¦  «         G d6„ d7e
j7        ¦  «        ¦   «         Z@d8„ ZA	 	 	 d‚d:e
j7        d;ej        d<ej        d=ej        d>ej        dz  d?e=e9z  d@e=dz  dAe=dz  d,eBej        ej        f         fdB„ZCdCej        dDej        dEej        dFej        d,eBej        ej        f         f
dG„ZDdHej        dIe9d,ej        fdJ„ZE G dK„ dLe
j7        ¦  «        ZF G dM„ dNe
j7        ¦  «        ZG G dO„ dPe
j7        ¦  «        ZH G dQ„ dRe
j7        ¦  «        ZI G dS„ dTe
j7        ¦  «        ZJ G dU„ dVe¦  «        ZK G dW„ dXe
j7        ¦  «        ZL G dY„ dZe
j7        ¦  «        ZM G d[„ d\e
j7        ¦  «        ZN G d]„ d^e
j7        ¦  «        ZO G d_„ d`e
j7        ¦  «        ZPe) G da„ dbe"¦  «        ¦   «         ZQ G dc„ ddeQ¦  «        ZRe) G de„ dfeQ¦  «        ¦   «         ZSe) G dg„ dheeQ¦  «        ¦   «         ZT e)di¬j¦  «         G dk„ dleQ¦  «        ¦   «         ZUdƒdn„ZV e)dodp¬q¦  «         G dr„ dseQ¦  «        ¦   «         ZW e)dtdu¬q¦  «         G dv„ dweQ¦  «        ¦   «         ZX e)dxdy¬q¦  «         G dz„ d{eQ¦  «        ¦   «         ZY e)d|d}¬q¦  «         G d~„ deQ¦  «        ¦   «         ZZg d€¢Z[dS )„é    N)ÚCallableÚIterable)Ú	dataclass)ÚTensorÚnné   )Úinitialization)ÚACT2FN)ÚBackboneMixinÚfilter_output_hidden_states)Úuse_kernel_forward_from_hub)ÚGradientCheckpointingLayer)ÚBackboneOutputÚBaseModelOutputÚBaseModelOutputWithPoolingÚModelOutputÚSemanticSegmenterOutput)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)Ú#compile_compatible_method_lru_cache)ÚTransformersKwargsÚauto_docstring)Úcan_return_tupleÚmaybe_autocastÚmerge_with_config_defaults)Úcapture_outputsé   )ÚSapiens2Configz·
    Output type of [`Sapiens2Backbone`], extending [`BackboneOutput`] with optional CLS tokens from
    each selected feature stage (used when `config.return_class_token=True`).
    )Úcustom_introc                   ó>   — e Zd ZU dZdZeej                 dz  ed<   dS )ÚSapiens2BackboneOutputzÙ
    cls_tokens (`tuple(torch.FloatTensor)`, *optional*):
        CLS token from each selected feature stage, each of shape `(batch_size, hidden_size)`.
        Only present when `config.return_class_token=True`.
    NÚ
cls_tokens)	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r#   ÚtupleÚtorchÚFloatTensorÚ__annotations__© ó    úl/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/sapiens2/modeling_sapiens2.pyr"   r"   1   s;   € € € € € € ðð ð 37€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6Ð6Ð6r-   r"   z6
    Class for outputs of pose estimation models.
    c                   ó¬   — 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
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )ÚSapiens2PoseEstimatorOutputaÔ  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Pose estimation loss.
    heatmaps (`torch.FloatTensor` of shape `(batch_size, num_keypoints, height, width)`):
        Heatmaps as predicted by the model.
    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, if the model has an embedding layer, +
        one for the output of each stage) of shape `(batch_size, sequence_length, hidden_size)`. Hidden-states
        (also called feature maps) of the model at the output of each stage.
    NÚlossÚheatmaps.Úhidden_statesÚ
attentions)r$   r%   r&   r'   r1   r)   r*   r+   r2   r3   r(   r4   r,   r-   r.   r0   r0   E   s’   € € € € € € ð	ð 	ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø)-€HˆeÔ $Ñ&Ð-Ð-Ñ-Ø:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>Ø7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r-   r0   z8
    Class for outputs of normal estimation models.
    c                   ó¬   — 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
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )ÚSapiens2NormalEstimatorOutputa  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Normal estimation loss.
    normals (`torch.FloatTensor` of shape `(batch_size, num_labels, height, width)`):
        Raw normal map predictions as output by the model (unnormalized).
    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 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 per layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`. Attentions weights after the attention softmax.
    Nr1   Únormals.r3   r4   )r$   r%   r&   r'   r1   r)   r*   r+   r7   r3   r(   r4   r,   r-   r.   r6   r6   ]   s’   € € € € € € ðð ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø(,€GˆUÔ Ñ%Ð,Ð,Ñ,Ø:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>Ø7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r-   r6   z:
    Class for outputs of pointmap estimation models.
    c                   óÊ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
eej        df         dz  ed<   dZeej        df         dz  ed<   dS )	ÚSapiens2PointmapEstimatorOutputaÄ  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Pointmap estimation loss.
    pointmaps (`torch.FloatTensor` of shape `(batch_size, 3, height, width)`):
        Per-pixel 3D XYZ coordinate predictions in canonical camera space.
    scales (`torch.FloatTensor` of shape `(batch_size, 1)`, *optional*):
        Canonical focal length / actual focal length ratio. `None` when no scale branch is configured.
    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 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 per layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`. Attentions weights after the attention softmax.
    Nr1   Ú	pointmapsÚscales.r3   r4   )r$   r%   r&   r'   r1   r)   r*   r+   r:   r;   r3   r(   r4   r,   r-   r.   r9   r9   x   sª   € € € € € € ðð ð  &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø*.€IˆuÔ  4Ñ'Ð.Ð.Ñ.Ø'+€FˆEÔ Ñ$Ð+Ð+Ñ+Ø:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>Ø7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r-   r9   z4
    Class for outputs of image matting models.
    c                   óÂ   — 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
ej                 dz  ed<   dZe
ej                 dz  ed<   dZej        dz  ed<   dS )ÚSapiens2ImageMattingOutputaN  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Loss.
    alphas (`torch.FloatTensor` of shape `(batch_size, 1, height, width)`):
        Estimated alpha values.
    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, if the model has an embedding layer, +
        one for the output of each stage) of shape `(batch_size, sequence_length, hidden_size)`. Hidden-states
        (also called feature maps) of the model at the output of each stage.
    foregrounds (`torch.FloatTensor` of shape `(batch_size, 3, height, width)`):
        Pre-multiplied RGB foreground predictions in `[0, 1]` (sigmoid-activated).
    Nr1   Úalphasr3   r4   Úforegrounds)r$   r%   r&   r'   r1   r)   r*   r+   r>   r3   r(   r4   r?   r,   r-   r.   r=   r=   –   s    € € € € € € ðð ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø'+€FˆEÔ Ñ$Ð+Ð+Ñ+Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø26€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6à,0€K�Ô" TÑ)Ð0Ð0Ñ0Ð0Ð0r-   r=   c                   ób   ‡ — e Zd ZdZdefˆ fd„Zd	dej        dej        dz  dej        fd„Zˆ xZ	S )
ÚSapiens2EmbeddingszM
    Construct the CLS token, mask token, position and patch embeddings.
    Úconfigc                 ó   •— t          ¦   «                              ¦   «          || _        t          j        t          j        dd|j        ¦  «        ¦  «        | _        |j	        r-t          j        t          j
        dd|j        ¦  «        ¦  «        nd | _        t          j        t          j        d|j        |j        ¦  «        ¦  «        | _        t          j        |j        |j        |j        |j        ¬¦  «        | _        d S )Nr   )Úkernel_sizeÚstride)ÚsuperÚ__init__rB   r   Ú	Parameterr)   ÚrandnÚhidden_sizeÚ	cls_tokenÚuse_mask_tokenÚzerosÚ
mask_tokenÚemptyÚnum_register_tokensÚregister_tokensÚConv2dÚnum_channelsÚ
patch_sizeÚpatch_embeddings©ÚselfrB   Ú	__class__s     €r.   rG   zSapiens2Embeddings.__init__·   sÎ   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝœ¥e¤k°!°Q¸Ô8JÑ&KÔ&KÑLÔLˆŒØQWÔQfÐp�"œ,¥u¤{°1°a¸Ô9KÑ'LÔ'LÑMÔMÐMÐlpˆŒÝ!œ|­E¬K¸¸6Ô;UÐW]ÔWiÑ,jÔ,jÑkÔkˆÔÝ "¤	ØÔ Ô!3ÀÔARÐ[aÔ[lð!
ñ !
ô !
ˆÔÐÐr-   NÚpixel_valuesÚbool_masked_posÚreturnc                 óL  — |�| j         €t          d¦  «        ‚|j        d         }| j        j        j        }|                      |                     |¬¦  «        ¦  «        }|                     d¦  «                             dd¦  «        }|�H| j                              |j        ¦  «        }t          j
        |                     d¦  «        ||¦  «        }| j                             |dd¦  «        }| j                             |dd¦  «        }t          j        |||gd¬¦  «        }	|	S )Nz:bool_masked_pos requires use_mask_token=True in the configr   ©Údtypeé   r   éÿÿÿÿ©Údim)rN   Ú
ValueErrorÚshaperU   Úweightr^   ÚtoÚflattenÚ	transposer)   ÚwhereÚ	unsqueezerK   ÚexpandrQ   Úcat)
rW   rY   rZ   Ú
batch_sizeÚtarget_dtyperU   rN   rK   rQ   Ú
embeddingss
             r.   ÚforwardzSapiens2Embeddings.forwardÁ   s  € ØÐ&¨4¬?Ð+BÝÐYÑZÔZÐZØ!Ô'¨Ô*ˆ
ØÔ,Ô3Ô9ˆð  ×0Ò0°·²À|°Ñ1TÔ1TÑUÔUÐØ+×3Ò3°AÑ6Ô6×@Ò@ÀÀAÑFÔFÐàÐ&Øœ×+Ò+Ð,<Ô,BÑCÔCˆJÝ$œ{¨?×+DÒ+DÀRÑ+HÔ+HÈ*ÐVfÑgÔgÐð ”N×)Ò)¨*°b¸"Ñ=Ô=ˆ	ØÔ.×5Ò5°jÀ"ÀbÑIÔIˆÝ”Y 	¨?Ð<LÐMÐSTÐUÑUÔUˆ
àÐr-   ©N)
r$   r%   r&   r'   r   rG   r)   r   rp   Ú__classcell__©rX   s   @r.   rA   rA   ²   sŠ   ø€ € € € € ðð ð
˜~ð 
ð 
ð 
ð 
ð 
ð 
ðð  E¤Lð À5Ä<ÐRVÑCVð ÐbgÔbnð ð ð ð ð ð ð ð r-   rA   é    )ÚmaxsizeÚnum_patches_hÚnum_patches_wr^   Údevicer[   c                 ó  — t          j        d| ||¬¦  «        }t          j        d|||¬¦  «        }|| z  }||z  }t          j        t          j        ||d¬¦  «        d¬¦  «        }|                     dd¦  «        }d	|z  d
z
  }|S )aq  
    Computes the 2D coordinates of the centers of image patches, normalized to the range [-1, +1].
    The center of each patch is exactly halfway between its top-left and bottom-right corners.

    Args:
        num_patches_h (int): Number of patches along the vertical (height) axis.
        num_patches_w (int): Number of patches along the horizontal (width) axis.
        dtype (torch.dtype): The desired data type of the returned tensor.

    Returns:
        torch.Tensor: A tensor of shape (height * width, 2), where each row contains the (y, x)
            coordinates of a patch center, normalized to [-1, +1].
    g      à?©r^   rx   Úij)Úindexingr`   ra   r   r   g       @g      ð?)r)   ÚarangeÚstackÚmeshgridrg   )rv   rw   r^   rx   Úcoords_hÚcoords_wÚcoordss          r.   Úget_patches_center_coordinatesrƒ   ×   s”   € õ" Œ|˜C °eÀFÐKÑKÔK€HÝŒ|˜C °eÀFÐKÑKÔK€HØ˜-Ñ'€HØ˜-Ñ'€HåŒ[�œ¨°(ÀTÐJÑJÔJÐPRÐSÑSÔS€FØ�^Š^˜A˜qÑ!Ô!€Fà�6‰\˜CÑ€FØ€Mr-   r‚   ÚshiftÚjitterÚrescalec                 ó  — |�=t          j        d| j        | j        ¬¦  «        }|                     | |¦  «        }| |z   } |�ct          j        |¦  «        }t          j        d| j        | j        ¬¦  «        }|                     | |¦  «                             ¦   «         }| |z  } |�ct          j        |¦  «        }t          j        d| j        | j        ¬¦  «        }|                     | |¦  «                             ¦   «         }| |z  } | S )N)r   r_   )rx   r^   r   )r)   rO   rx   r^   Úuniform_ÚnpÚlogÚexp)	r‚   r„   r…   r†   Úshift_hwÚjitter_rangeÚ	jitter_hwÚrescale_rangeÚ
rescale_hws	            r.   Ú"augment_patches_center_coordinatesr‘   ô   s  € ð ÐÝ”;˜v¨f¬mÀ6Ä<ÐPÑPÔPˆØ×$Ò$ e V¨UÑ3Ô3ˆØ˜(Ñ"ˆð ÐÝ”v˜f‘~”~ˆÝ”K ¨v¬}ÀFÄLÐQÑQÔQˆ	Ø×&Ò&¨ }°lÑCÔC×GÒGÑIÔIˆ	Ø˜)Ñ#ˆð ÐÝœ˜w™œˆÝ”[ ¨6¬=ÀÄÐMÑMÔMˆ
Ø×(Ò(¨-¨¸ÑGÔG×KÒKÑMÔMˆ
Ø˜*Ñ$ˆà€Mr-   c                   óx   ‡ — e Zd ZU ej        ed<   defˆ fd„Zdej        deej        ej        f         fd„Z	ˆ xZ
S )ÚSapiens2RopePositionEmbeddingÚinv_freqrB   c                 ó,  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        |j        z  | _        d| j        t          j	        ddd| j        z  t          j
        ¬¦  «        z  z  }|                      d|d¬¦  «         |j        }t          |t          ¦  «        r|n||f\  }}|j        }t          |t           ¦  «        r|n|d         }t          |t           ¦  «        r|n|d         }||z  | _        ||z  | _        d S )Nr   r   é   r]   r”   F)Ú
persistent)rF   rG   rB   Ú
rope_thetaÚbaserJ   Únum_attention_headsÚhead_dimr)   r}   Úfloat32Úregister_bufferÚ
image_sizeÚ
isinstancer   rT   Úintrv   rw   )
rW   rB   r”   rž   Úimage_hÚimage_wrT   Úpatch_size_hÚpatch_size_wrX   s
            €r.   rG   z&Sapiens2RopePositionEmbedding.__init__  s  ø€ Ý‰Œ×ÒÑÔÐàˆŒØÔ%ˆŒ	ØÔ*¨fÔ.HÑHˆŒà�t”y¥E¤L°°A°q¸4¼=Ñ7HÕPUÔP]Ð$^Ñ$^Ô$^Ñ^Ñ^ˆØ×Ò˜Z¨¸eÐÑDÔDÐDØÔ&ˆ
Ý)3°JÅÑ)IÔ)IÐg˜:˜:ÐPZÐ\fÐOgÑˆ�ØÔ&ˆ
Ý%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆÝ%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆØ$¨Ñ4ˆÔØ$¨Ñ4ˆÔÐÐr-   rY   r[   c                 ó,  — |j         \  }}}}|| j        j        z  }|| j        j        z  }|j        }t	          |j        t          ¦  «        r|j        dk    r|j        nd}t          |d¬¦  «        5  t          ||t          j
        |¬¦  «        }	| j        r1t          |	| j        j        | j        j        | j        j        ¬¦  «        }	dt           j        z  |	d d …d d …d f         z  | j        d d d d …f         z  }
|
                     dd¦  «        }
|
                     d¦  «        }
t          j        |
¦  «        }t          j        |
¦  «        }d d d ¦  «         n# 1 swxY w Y   |j        }|                     |¬	¦  «        |                     |¬	¦  «        fS )
NÚmpsÚcpuF)Údevice_typeÚenabledrz   )r„   r…   r†   r_   r   r]   )rd   rB   rT   rx   rŸ   ÚtypeÚstrr   rƒ   r)   rœ   Útrainingr‘   Úpos_embed_shiftÚpos_embed_jitterÚpos_embed_rescaleÚmathÚpir”   rg   ÚtileÚcosÚsinr^   rf   )rW   rY   Ú_ÚheightÚwidthrv   rw   rx   r¨   Úpatch_coordsÚanglesr³   r´   r^   s                 r.   rp   z%Sapiens2RopePositionEmbedding.forward%  s×  € Ø*Ô0Ñˆˆ1ˆf�eØ $¤+Ô"8Ñ8ˆØ ¤Ô!7Ñ7ˆàÔ$ˆÝ%/°´½SÑ%AÔ%AÐeÀfÄkÐUZÒFZÐFZ�f”k�kÐ`eˆå¨¸UÐCÑCÔCð 	$ð 	$õ :Ø˜}µE´MÈ&ðñ ô ˆLð Œ}ð ÝAØ Øœ+Ô5Øœ;Ô7Ø œKÔ9ð	 ñ  ô  �ð �œ‘[ <°°°°1°1°1°d°
Ô#;Ñ;¸d¼mÈDÐRVÐXYÐXYÐXYÈMÔ>ZÑZˆFØ—^’^ A qÑ)Ô)ˆFØ—[’[ ‘^”^ˆFå”)˜FÑ#Ô#ˆCÝ”)˜FÑ#Ô#ˆCð+	$ð 	$ð 	$ñ 	$ô 	$ð 	$ð 	$ð 	$ð 	$ð 	$ð 	$øøøð 	$ð 	$ð 	$ð 	$ð. Ô"ˆØ�vŠv˜EˆvÑ"Ô" C§F¢F° FÑ$7Ô$7Ð7Ð7s   Á1CEÅEÅE)r$   r%   r&   r)   r   r+   r   rG   r(   rp   rr   rs   s   @r.   r“   r“     s†   ø€ € € € € € ØŒlÐÐÑð5˜~ð 5ð 5ð 5ð 5ð 5ð 5ð" 8 E¤Lð  8°U¸5¼<ÈÌÐ;UÔ5Vð  8ð  8ð  8ð  8ð  8ð  8ð  8ð  8r-   r“   ÚRMSNormc                   óT   ‡ — e Zd Zd	deddfˆ fd„Zdej        dej        fd„Zd„ Zˆ xZ	S )
ÚSapiens2RMSNormç�íµ ÷Æ°>Úepsr[   Nc                 ó¬   •— t          ¦   «                              ¦   «          t          j        t	          j        |¦  «        ¦  «        | _        || _        dS )z>
        Sapiens2RMSNorm is equivalent to T5LayerNorm
        N)rF   rG   r   rH   r)   Úonesre   Úvariance_epsilon)rW   rJ   r¾   rX   s      €r.   rG   zSapiens2RMSNorm.__init__J  sD   ø€ õ 	‰Œ×ÒÑÔÐÝ”l¥5¤:¨kÑ#:Ô#:Ñ;Ô;ˆŒØ #ˆÔÐÐr-   r3   c                 ó  — |j         }|                     t          j        ¦  «        }|                     d¦  «                             dd¬¦  «        }|t          j        || j        z   ¦  «        z  }| j        |                     |¦  «        z  S )Nr_   r`   T)Úkeepdim)	r^   rf   r)   rœ   ÚpowÚmeanÚrsqrtrÁ   re   )rW   r3   Úinput_dtypeÚvariances       r.   rp   zSapiens2RMSNorm.forwardR  s|   € Ø#Ô)ˆØ%×(Ò(­¬Ñ7Ô7ˆØ ×$Ò$ QÑ'Ô'×,Ò,¨R¸Ð,Ñ>Ô>ˆØ%­¬°H¸tÔ?TÑ4TÑ(UÔ(UÑUˆØŒ{˜]×-Ò-¨kÑ:Ô:Ñ:Ð:r-   c                 óH   — t          | j        j        ¦  «        › d| j        › �S )Nz, eps=)r(   re   rd   rÁ   ©rW   s    r.   Ú
extra_reprzSapiens2RMSNorm.extra_reprY  s&   € Ý˜œÔ)Ñ*Ô*ÐIÐI°$Ô2GÐIÐIÐIr-   )r½   )
r$   r%   r&   ÚfloatrG   r)   r   rp   rË   rr   rs   s   @r.   r¼   r¼   H  sŒ   ø€ € € € € ð$ð $¨ð $¸$ð $ð $ð $ð $ð $ð $ð; U¤\ð ;°e´lð ;ð ;ð ;ð ;ðJð Jð Jð Jð Jð Jð Jr-   r¼   c                 óœ   — | dd| j         d         dz  …f         }| d| j         d         dz  d…f         }t          j        | |fd¬¦  «        S )z*Rotates half the hidden dims of the input..Nr`   r_   ra   )rd   r)   rl   )ÚxÚx1Úx2s      r.   Úrotate_halfrÑ   ]  s]   € à	
ˆ3Ð"�!”'˜"”+ Ñ"Ð"Ð"Ô	#€BØ	
ˆ3�”˜”˜qÑ Ð"Ð"Ð"Ô	#€BÝŒ9�r�c˜2�Y BÐ'Ñ'Ô'Ð'r-   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚdropoutÚscalingÚsoftcapc                 ól  — |€
| j         dz  }t          || j        ¦  «        }	t          || j        ¦  «        }
t          j        ||	                     dd¦  «        ¦  «        |z  }|�||z  }t          j        |¦  «        }||z  }|�||z   }t          j         	                    |dt          j
        ¬¦  «                             |j        ¦  «        }t          j                             ||| j        ¬¦  «        }t          j        ||
¦  «        }|                     dd¦  «                             ¦   «         }||fS )Nç      à¿r_   r   r`   )rb   r^   )Úpr¬   r   )r›   Ú	repeat_kvÚnum_key_value_groupsr)   Úmatmulrh   Útanhr   Ú
functionalÚsoftmaxrœ   rf   r^   rØ   r¬   Ú
contiguous)rÓ   rÔ   rÕ   rÖ   r×   rØ   rÙ   rÚ   ÚkwargsÚ
key_statesÚvalue_statesÚattn_weightsÚattn_outputs                r.   Úeager_attention_forwardrê   d  s%  € ð €Ø”/ 4Ñ'ˆå˜3 Ô ;Ñ<Ô<€JÝ˜U FÔ$?Ñ@Ô@€Lå”<  z×';Ò';¸A¸qÑ'AÔ'AÑBÔBÀWÑL€LàÐØ# gÑ-ˆÝ”z ,Ñ/Ô/ˆØ# gÑ-ˆØÐ!Ø# nÑ4ˆõ ”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐW\ÔWbÑcÔc€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€LÝ”,˜|¨\Ñ:Ô:€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€KØ˜Ð$Ð$r-   ÚqÚkr³   r´   c                 óx  — | j         d         }|j         d         }||z
  }|                      ||fd¬¦  «        \  }}	|                     ||fd¬¦  «        \  }
}|	|z  t          |	¦  «        |z  z   }	||z  t          |¦  «        |z  z   }t          j        ||	fd¬¦  «        } t          j        |
|fd¬¦  «        }| |fS )a  Applies Rotary Position Embedding to the query and key tensors, but only to the patch tokens,
    ignoring the prefix tokens (cls token and register tokens).

    Args:
        q (`torch.Tensor`): The query tensor.
        k (`torch.Tensor`): The key tensor.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.

    Returns:
        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
    éþÿÿÿra   )rd   ÚsplitrÑ   r)   rl   )rë   rì   r³   r´   rå   Ú
num_tokensÚnum_patchesÚnum_prefix_tokensÚq_prefix_tokensÚ	q_patchesÚk_prefix_tokensÚ	k_patchess               r.   Úapply_rotary_pos_embr÷   †  sØ   € ð  ”˜”€JØ”)˜B”-€KØ" [Ñ0Ðà!"§¢Ð*;¸[Ð)IÈr Ñ!RÔ!RÑ€O�YØ!"§¢Ð*;¸[Ð)IÈr Ñ!RÔ!RÑ€O�Yð ˜S‘¥[°Ñ%;Ô%;¸cÑ%AÑB€IØ˜S‘¥[°Ñ%;Ô%;¸cÑ%AÑB€IåŒ	�? IÐ.°BÐ7Ñ7Ô7€AÝŒ	�? IÐ.°BÐ7Ñ7Ô7€Aàˆaˆ4€Kr-   r3   Ún_repc                 ó¸   — | j         \  }}}}|dk    r| S | dd…dd…ddd…dd…f                              |||||¦  «        } |                      |||z  ||¦  «        S )zÔ
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    r   N)rd   rk   Úreshape)r3   rø   ÚbatchÚnum_key_value_headsÚslenr›   s         r.   rÞ   rÞ   §  s„   € ð
 2?Ô1DÑ.€EÐ  hØ�‚z€zØÐØ! ! ! ! Q Q Q¨¨a¨a¨a°°°Ð"2Ô3×:Ò:¸5ÐBUÐW\Ð^bÐdlÑmÔm€MØ× Ò  Ð(;¸eÑ(CÀTÈ8ÑTÔTÐTr-   c                   óÈ   ‡ — e Zd ZdZdedefˆ fd„Z	 	 ddej        dej        dz  de	ej        ej        f         dz  d	e
e         d
e	ej        ej        dz  f         f
d„Zˆ xZS )ÚSapiens2AttentionzI
    Multi-headed attention compatible with ALL_ATTENTION_FUNCTIONS.
    rB   Ú	layer_idxc                 ó¬  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        | j        | j        z  | _        d| _        | j        dz  | _	        d| _        |j
        | _        t          j        | j        | j        |j        ¬¦  «        | _        t          j        | j        | j        |j        ¬¦  «        | _        |j        |         | _        | j        | j        z  | _        t          j        | j        | j        | j        z  |j        ¬¦  «        | _        t          j        | j        | j        | j        z  |j        ¬¦  «        | _        |j        rt5          | j        |j        ¬¦  «        nt          j        ¦   «         | _        |j        rt5          | j        |j        ¬¦  «        nt          j        ¦   «         | _        d S )NFrÜ   ©Úbias©r¾   )rF   rG   rB   rJ   Ú	embed_dimrš   Ú	num_headsr›   Ú	is_causalrÙ   Úattention_dropoutrØ   r   ÚLinearÚ
query_biasÚq_projÚ	proj_biasÚo_projÚnum_key_value_heads_per_layerrü   rß   Úkey_biasÚk_projÚ
value_biasÚv_projÚuse_qk_normr¼   Úrms_norm_epsÚIdentityÚq_normÚk_norm©rW   rB   r   rX   s      €r.   rG   zSapiens2Attention.__init__¸  s~  ø€ Ý‰Œ×ÒÑÔÐØˆŒØÔ+ˆŒØÔ3ˆŒØœ¨$¬.Ñ8ˆŒØˆŒà”} dÑ*ˆŒØˆŒàÔ/ˆŒå”i ¤°´ÀVÔEVÐWÑWÔWˆŒÝ”i ¤°´ÀVÔEUÐVÑVÔVˆŒØ#)Ô#GÈ	Ô#RˆÔ Ø$(¤N°dÔ6NÑ$NˆÔ!Ý”i ¤°Ô0HÈ4Ì=Ñ0XÐ_eÔ_nÐoÑoÔoˆŒÝ”i ¤°Ô0HÈ4Ì=Ñ0XÐ_eÔ_pÐqÑqÔqˆŒØQWÔQcÐv•o d¤m¸Ô9LÐMÑMÔMÐMÕikÔitÑivÔivˆŒØQWÔQcÐv•o d¤m¸Ô9LÐMÑMÔMÐMÕikÔitÑivÔivˆŒˆˆr-   Nr3   r×   Úposition_embeddingsrå   r[   c                 ó4  — |j         dd…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }	|                      |¦  «        }|                      |¦  «        }|\  }
}t          |||
|¦  «        \  }}t          j        | j        j        t          ¦  «        } || |||	|f| j        sdn| j        | j        dœ|¤Ž\  }} |j        g |¢d‘R Ž                      ¦   «         }|                      |¦  «        }||fS )z#Input shape: Batch x Time x ChannelNr`   r   r_   rÒ   )rØ   rÙ   )rd   r›   r  Úviewrh   r  r  r  r  r÷   r   Úget_interfacerB   Ú_attn_implementationrê   r¬   rØ   rÙ   rú   rä   r  )rW   r3   r×   r  rå   Úinput_shapeÚhidden_shapeÚquery_statesræ   rç   r³   r´   Úattention_interfaceré   rè   s                  r.   rp   zSapiens2Attention.forwardÎ  s»  € ð $Ô)¨#¨2¨#Ô.ˆØ8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆà—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆØ—[’[ Ñ/Ô/×4Ò4°\ÑBÔB×LÒLÈQÐPQÑRÔRˆ
Ø—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆØ—{’{ <Ñ0Ô0ˆØ—[’[ Ñ,Ô,ˆ
à&‰ˆˆSÝ#7¸ÀjÐRUÐWZÑ#[Ô#[Ñ ˆ�jå(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØð	%
ð  $œ}Ð>�C�C°$´,Ø”Lð	%
ð 	%
ð ð	%
ð 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—k’k +Ñ.Ô.ˆØ˜LÐ(Ð(r-   ©NN©r$   r%   r&   r'   r   r    rG   r)   r   r(   r   r   rp   rr   rs   s   @r.   rÿ   rÿ   ³  sã   ø€ € € € € ðð ðw˜~ð w¸#ð wð wð wð wð wð wð2 /3ØHLð	$)ð $)à”|ð$)ð œ tÑ+ð$)ð # 5¤<°´Ð#=Ô>ÀÑEð	$)ð
 Ð+Ô,ð$)ð 
ˆuŒ|˜Uœ\¨DÑ0Ð0Ô	1ð$)ð $)ð $)ð $)ð $)ð $)ð $)ð $)r-   rÿ   c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )ÚSapiens2LayerScaler[   Nc                 ó¸   •— t          ¦   «                              ¦   «          t          j        |j        t          j        |j        ¦  «        z  ¦  «        | _        d S rq   )	rF   rG   r   rH   Úlayerscale_valuer)   rÀ   rJ   Úlambda1rV   s     €r.   rG   zSapiens2LayerScale.__init__ö  sC   ø€ Ý‰Œ×ÒÑÔÐÝ”| FÔ$;½e¼jÈÔI[Ñ>\Ô>\Ñ$\Ñ]Ô]ˆŒˆˆr-   Úhidden_statec                 ó   — || j         z  S rq   )r(  )rW   r)  s     r.   rp   zSapiens2LayerScale.forwardú  s   € Ø˜dœlÑ*Ð*r-   ©r[   N)r$   r%   r&   rG   r)   r   rp   rr   rs   s   @r.   r%  r%  õ  si   ø€ € € € € ð^ð ^ð ^ð ^ð ^ð ^ð+ E¤Lð +°U´\ð +ð +ð +ð +ð +ð +ð +ð +r-   r%  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚSapiens2MLPc                 ó`  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        t          j        | j        | j        |j        ¬¦  «        | _        t          j        | j        | j        |j        ¬¦  «        | _	        t          |j                 | _        d S ©Nr  )rF   rG   rB   rJ   Úintermediate_sizer   r	  Úmlp_biasÚup_projÚ	down_projr
   Ú
hidden_actÚact_fnrV   s     €r.   rG   zSapiens2MLP.__init__ÿ  s�   ø€ Ý‰Œ×ÒÑÔÐØˆŒØ!Ô-ˆÔØ!'Ô!9ˆÔÝ”y Ô!1°4Ô3IÐPVÔP_Ð`Ñ`Ô`ˆŒÝœ 4Ô#9¸4Ô;KÐRXÔRaÐbÑbÔbˆŒÝ˜VÔ.Ô/ˆŒˆˆr-   c                 óx   — |                       |                      |                      |¦  «        ¦  «        ¦  «        S rq   )r3  r5  r2  )rW   rÎ   s     r.   rp   zSapiens2MLP.forward  s*   € Ø�~Š~˜dŸkšk¨$¯,ª,°q©/¬/Ñ:Ô:Ñ;Ô;Ð;r-   ©r$   r%   r&   rG   rp   rr   rs   s   @r.   r-  r-  þ  sG   ø€ € € € € ð0ð 0ð 0ð 0ð 0ð<ð <ð <ð <ð <ð <ð <r-   r-  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚSapiens2GatedMLPc                 ó¶  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        t          j        | j        | j        |j        ¬¦  «        | _        t          j        | j        | j        |j        ¬¦  «        | _	        t          j        | j        | j        |j        ¬¦  «        | _
        t          |j                 | _        d S r/  )rF   rG   rB   rJ   r0  r   r	  r1  Ú	gate_projr2  r3  r
   r4  r5  rV   s     €r.   rG   zSapiens2GatedMLP.__init__  s¯   ø€ Ý‰Œ×ÒÑÔÐØˆŒØ!Ô-ˆÔØ!'Ô!9ˆÔÝœ 4Ô#3°TÔ5KÐRXÔRaÐbÑbÔbˆŒÝ”y Ô!1°4Ô3IÐPVÔP_Ð`Ñ`Ô`ˆŒÝœ 4Ô#9¸4Ô;KÐRXÔRaÐbÑbÔbˆŒÝ˜VÔ.Ô/ˆŒˆˆr-   c                 ó¨   — |                       |                      |                      |¦  «        ¦  «        |                      |¦  «        z  ¦  «        }|S rq   )r3  r5  r;  r2  )rW   rÎ   r3  s      r.   rp   zSapiens2GatedMLP.forward  sA   € Ø—N’N 4§;¢;¨t¯~ª~¸aÑ/@Ô/@Ñ#AÔ#AÀDÇLÂLÐQRÁOÄOÑ#SÑTÔTˆ	ØÐr-   r7  rs   s   @r.   r9  r9    sG   ø€ € € € € ð0ð 0ð 0ð 0ð 0ðð ð ð ð ð ð r-   r9  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 )ÚSapiens2DropPathzÏ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_probr[   Nc                 óV   •— t          ¦   «                              ¦   «          || _        d S rq   )rF   rG   r?  )rW   r?  rX   s     €r.   rG   zSapiens2DropPath.__init__#  s$   ø€ Ý‰Œ×ÒÑÔÐØ"ˆŒˆˆr-   r3   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   rz   )
r?  r¬   rd   Úndimr)   Úrandr^   rx   ÚfloorÚdiv)rW   r3   Ú	keep_probrd   Úrandom_tensors        r.   rp   zSapiens2DropPath.forward'  s“   € ØŒ>˜SÒ Ð ¨¬Ð Ø Ð Ø˜œÑ&ˆ	ØÔ$ QÔ'Ð)¨D°MÔ4FÈÑ4JÑ,KÑKˆÝœ
 5°Ô0CÈMÔL`ÐaÑaÔaˆÝœ M°IÑ$=Ñ>Ô>ˆØ× Ò  Ñ+Ô+¨mÑ;Ð;r-   c                 ó   — d| j         › �S )Nzp=)r?  rÊ   s    r.   rË   zSapiens2DropPath.extra_repr0  s   € Ø$�D”NÐ$Ð$Ð$r-   )rÒ   )r$   r%   r&   r'   rÌ   rG   r)   r   rp   r«   rË   rr   rs   s   @r.   r>  r>    s›   ø€ € € € € ðð ð#ð # %ð #°$ð #ð #ð #ð #ð #ð #ð< U¤\ð <°e´lð <ð <ð <ð <ð%˜Cð %ð %ð %ð %ð %ð %ð %ð %r-   r>  c                   ó¨   ‡ — e Zd ZdZdedefˆ fd„Z	 	 ddej        dej        dz  de	ej        ej        f         dz  d	e
e         d
ej        f
d„Zˆ xZS )ÚSapiens2LayerzCThis corresponds to the Block class in the original implementation.rB   r   c                 ó  •— t          ¦   «                              ¦   «          t          |j        |j        ¬¦  «        | _        t          ||¬¦  «        | _        t          |¦  «        | _	        |j
        dk    rt          |j
        ¦  «        nt          j        ¦   «         | _        t          |j        |j        ¬¦  «        | _        |j        rt#          |¦  «        | _        nt'          |¦  «        | _        t          j        ¦   «         | _        d S )Nr  ©r   rÒ   )rF   rG   r¼   rJ   r  Únorm1rÿ   Ú	attentionr%  Úlayer_scale1Údrop_path_rater>  r   r  Ú	drop_pathÚnorm2Úuse_gated_mlpr9  Úmlpr-  Úlayer_scale2r  s      €r.   rG   zSapiens2Layer.__init__7  sà   ø€ Ý‰Œ×ÒÑÔÐÝ$ VÔ%7¸VÔ=PÐQÑQÔQˆŒ
Ý*¨6¸YÐGÑGÔGˆŒÝ.¨vÑ6Ô6ˆÔØDJÔDYÐ\_ÒD_ÐD_Õ)¨&Ô*?Ñ@Ô@Ð@ÕegÔepÑerÔerˆŒÝ$ VÔ%7¸VÔ=PÐQÑQÔQˆŒ
àÔð 	+Ý'¨Ñ/Ô/ˆDŒHˆHå" 6Ñ*Ô*ˆDŒHÝœK™MœMˆÔÐÐr-   Nr3   r×   r  rå   r[   c                 óh  — |}|                       |¦  «        } | j        |f||dœ|¤Ž\  }}|                      |¦  «        }|                      |¦  «        |z   }|}|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        |z   }|S )N)r×   r  )rM  rN  rO  rQ  rR  rT  rU  )rW   r3   r×   r  rå   Úresidualrµ   s          r.   rp   zSapiens2Layer.forwardE  sÒ   € ð !ˆØŸ
š
 =Ñ1Ô1ˆØ)˜4œ>Øð
à)Ø 3ð
ð 
ð ð	
ð 
Ñˆ�qð ×)Ò)¨-Ñ8Ô8ˆØŸš }Ñ5Ô5¸Ñ@ˆð !ˆØŸ
š
 =Ñ1Ô1ˆØŸš Ñ/Ô/ˆØ×)Ò)¨-Ñ8Ô8ˆØŸš }Ñ5Ô5¸Ñ@ˆàÐr-   r"  r#  rs   s   @r.   rJ  rJ  4  sÆ   ø€ € € € € ØMÐMð*˜~ð *¸#ð *ð *ð *ð *ð *ð *ð" /3ØHLð	ð à”|ðð œ tÑ+ðð # 5¤<°´Ð#=Ô>ÀÑEð	ð
 Ð+Ô,ðð 
Œðð ð ð ð ð ð ð r-   rJ  c                   óº   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 dded	ed
eeeef         z  dedeeeef         z  ez  dedededededefˆ fd„Zde	j
        de	j
        fd„Zˆ xZS )ÚSapiens2ConvLayerzc
    A basic wrapper for Convolution-BatchNorm-Activation, typically used for head components.
    r   r   ÚsiluTFr_   Úin_channelsÚout_channelsrD   rE   ÚpaddingÚgroupsÚ
activationr  Úconvolution_transposeÚpixel_shuffleÚscale_factorc           	      ó(  •— t          ¦   «                              ¦   «          |	rt          j        ||||¬¦  «        | _        n t          j        |||||||¬¦  «        | _        t          j        |¦  «        | _        t          |         | _	        |	r+t          j        ||
r||dz  z  n||||||¬¦  «        | _        n*t          j        ||
r||dz  z  n||||||¬¦  «        | _        |
rt          j
        |¦  «        nt          j        ¦   «         | _        d S )N)r[  r\  rD   rE   )r[  r\  rD   rE   r]  r^  r  r_   )rD   rE   r]  r  r^  )rF   rG   r   ÚConvTranspose2dÚconvolutionrR   ÚInstanceNorm2dÚnormr
   r5  ÚPixelShuffler  ra  )rW   r[  r\  rD   rE   r]  r^  r_  r  r`  ra  rb  rX   s               €r.   rG   zSapiens2ConvLayer.__init__g  s[  ø€ õ 	‰Œ×ÒÑÔÐØ ð 	Ý!Ô1Ø'Ø)Ø'Øð	 ñ  ô  ˆDÔÐõ  "œyØ'Ø)Ø'ØØØØð ñ  ô  ˆDÔõ Ô% lÑ3Ô3ˆŒ	Ý˜ZÔ(ˆŒØ ð 	Ý!Ô1ØØ2?ÐQ�˜|¨Q™Ñ.Ð.À\Ø'ØØØØð ñ  ô  ˆDÔÐõ  "œyØØ2?ÐQ�˜|¨Q™Ñ.Ð.À\Ø'ØØØØð ñ  ô  ˆDÔð ?LÐ^�Rœ_¨\Ñ:Ô:Ð:ÕQSÔQ\ÑQ^ÔQ^ˆÔÐÐr-   r3   r[   c                 ó®   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|S rq   )re  ra  rg  r5  ©rW   r3   s     r.   rp   zSapiens2ConvLayer.forwardŸ  sP   € Ø×(Ò(¨Ñ7Ô7ˆØ×*Ò*¨=Ñ9Ô9ˆØŸ	š	 -Ñ0Ô0ˆØŸš MÑ2Ô2ˆØÐr-   )	r   r   r   r   rZ  TFFr_   )r$   r%   r&   r'   r    r(   r«   ÚboolrG   r)   r   rp   rr   rs   s   @r.   rY  rY  b  s-  ø€ € € € € ðð ð ./ØØ/0ØØ ØØ&+Ø#Øð6_ð 6_àð6_ð ð6_ð ˜5  c œ?Ñ*ð	6_ð
 ð6_ð �u˜S #˜X”Ñ&¨Ñ,ð6_ð ð6_ð ð6_ð ð6_ð  $ð6_ð ð6_ð ð6_ð 6_ð 6_ð 6_ð 6_ð 6_ðp U¤\ð °e´lð ð ð ð ð ð ð ð r-   rY  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚSapiens2HeadrB   c                 ó>  •‡— t          ¦   «                              ¦   «          ‰j        j        rt	          ‰j        ‰j        dd¬¦  «        nt          j        ¦   «         | _        ‰j        g‰j        j	        d d…         z   }t          j
        ˆfd„t          |‰j        j	        ‰j        j        ¦  «        D ¦   «         ¦  «        | _        ‰j        j	        d         g‰j        j        d d…         z   }t          j
        ˆfd„t          |‰j        j        ‰j        j        ¦  «        D ¦   «         ¦  «        | _        ‰j        j        r‰j        j        d         n$‰j        j	        r‰j        j	        d         n‰j        }t          j        |‰j        d¬¦  «        | _        d S )Nr   r   ©rD   r]  r`   c              3   ó  •K  — | ]z\  }}}t          |||‰j        j        rd nd‰j        j        r|d z
  dz  nd t          ‰j        j        ¦  «        t          ‰j        j        ¦  «        ‰j        j         ¬¦  «        V — Œ{dS )r   r_   )rD   rE   r]  r  ra  r`  N)rY  Úhead_configÚuse_pixel_shufflerk  ©Ú.0Úin_chÚout_chrD   rB   s       €r.   ú	<genexpr>z(Sapiens2Head.__init__.<locals>.<genexpr>°  s´   øè è € ð -
ð -
ñ +��v˜{õ ØØØ'Ø"Ô.Ô@ÐG�q�qÀaØ28Ô2DÔ2VÐ]˜ q™¨QÑ.Ð.Ð\]Ý˜&Ô,Ô>Ñ?Ô?Ý" 6Ô#5Ô#GÑHÔHØ*0Ô*<Ô*NÐ&Nð	ñ 	ô 	ð-
ð -
ð -
ð -
ð -
ð -
r-   c              3   ón   •K  — | ]/\  }}}t          |||‰j        j        r|d z
  dz  nd¬¦  «        V — Œ0dS )r   r_   r   ro  N)rY  rq  rr  rs  s       €r.   rw  z(Sapiens2Head.__init__.<locals>.<genexpr>Â  st   øè è € ð 
)
ð 
)
ñ +��v˜{õ ØØØ'Ø28Ô2DÔ2VÐ]˜ q™¨QÑ.Ð.Ð\]ð	ñ ô ð
)
ð 
)
ð 
)
ð 
)
ð 
)
ð 
)
r-   )rD   )rF   rG   rq  rr  rY  rJ   r   r  Ú
input_convÚupsample_out_channelsÚ
ModuleListÚzipÚupsample_kernel_sizesÚupsample_layersÚconv_out_channelsÚconv_kernel_sizesÚconv_layersrR   Ú
num_labelsÚ	predictor)rW   rB   Úupsample_in_channelsÚconv_in_channelsÚpredictor_inrX   s    `   €r.   rG   zSapiens2Head.__init__¨  sÌ  øø€ Ý‰Œ×ÒÑÔÐð Ô!Ô3ðÕ˜fÔ0°&Ô2DÐRSÐ]^Ð_Ñ_Ô_Ð_å”‘”ð 	Œð
 !'Ô 2Ð3°fÔ6HÔ6^Ð_bÐ`bÐ_bÔ6cÑcÐÝ!œ}ð -
ð -
ð -
ð -
õ /2Ø$ØÔ"Ô8ØÔ"Ô8ñ/ô /ð-
ñ -
ô -
ñ  
ô  
ˆÔð" #Ô.ÔDÀRÔHÐIÈFÔL^ÔLpÐqtÐrtÐqtÔLuÑuÐÝœ=ð 
)
ð 
)
ð 
)
ð 
)
õ /2Ø  &Ô"4Ô"FÈÔHZÔHlñ/ô /ð
)
ñ 
)
ô 
)
ñ 

ô 

ˆÔð Ô!Ô3ð$ˆFÔÔ0°Ô4Ð4ð Ô!Ô7ð$�Ô#Ô9¸"Ô=Ð=àÔ#ð 	õ œ <°Ô1BÐPQÐRÑRÔRˆŒˆˆr-   r3   r[   c                 óª   — |                       |¦  «        }| j        D ]} ||¦  «        }Œ| j        D ]} ||¦  «        }Œ|                      |¦  «        S rq   )ry  r~  r�  rƒ  ©rW   r3   Úlayers      r.   rp   zSapiens2Head.forwardÖ  sk   € ØŸš¨Ñ6Ô6ˆØÔ)ð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMØÔ%ð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMØ�~Š~˜mÑ,Ô,Ð,r-   ©	r$   r%   r&   r   rG   r)   r   rp   rr   rs   s   @r.   rm  rm  §  sr   ø€ € € € € ð,S˜~ð ,Sð ,Sð ,Sð ,Sð ,Sð ,Sð\- U¤\ð -°e´lð -ð -ð -ð -ð -ð -ð -ð -r-   rm  c                   óJ   ‡ — e Zd Zdededej        ddfˆ fd„Zdedefd„Zˆ xZ	S )	ÚSapiens2PointmapFinalLayerBlockÚin_dimÚout_dimr_  r[   Nc                 ó¤   •— t          ¦   «                              ¦   «          t          j        t          j        ||¦  «        |g¦  «        | _        d S rq   )rF   rG   r   r{  r	  Úlayers)rW   r�  rŽ  r_  rX   s       €r.   rG   z(Sapiens2PointmapFinalLayerBlock.__init__à  s?   ø€ Ý‰Œ×ÒÑÔÐÝ”m¥R¤Y¨v°wÑ%?Ô%?ÀÐ$LÑMÔMˆŒˆˆr-   Úinputc                 ó4   — |}| j         D ]} ||¦  «        }Œ|S rq   )r�  )rW   r‘  r)  r‰  s       r.   rp   z'Sapiens2PointmapFinalLayerBlock.forwardä  s/   € ØˆØ”[ð 	/ð 	/ˆEØ ˜5 Ñ.Ô.ˆLˆLØÐr-   )
r$   r%   r&   r    r   ÚModulerG   r   rp   rr   rs   s   @r.   rŒ  rŒ  ß  s‡   ø€ € € € € ðN˜sð N¨Sð N¸b¼ið NÈDð Nð Nð Nð Nð Nð Nð˜Vð ¨ð ð ð ð ð ð ð ð r-   rŒ  c            	       óf   ‡ — e Zd Zddedeeef         dedefˆ fd„Zdej        d	ej        fd
„Z	ˆ xZ
S )ÚSapiens2PointmapFinalLayerr   rZ  r�  Úhidden_sizesrŽ  r_  c                 ód  •— t          ¦   «                              ¦   «          t          j        ¦   «         | _        t          ||d         t          |         ¬¦  «        | _        t          |d         |d         t          |         ¬¦  «        | _        t          j	        |d         |¦  «        | _
        d S )Nr   )r�  rŽ  r_  r   )rF   rG   r   ÚFlattenrg   rŒ  r
   Úblock1Úblock2r	  Úproj)rW   r�  r–  rŽ  r_  rX   s        €r.   rG   z#Sapiens2PointmapFinalLayer.__init__ì  s•   ø€ Ý‰Œ×ÒÑÔÐÝ”z‘|”|ˆŒÝ5Ø <°¤?½vÀjÔ?Qð
ñ 
ô 
ˆŒõ 6Ø ”?¨L¸¬OÍÈzÔHZð
ñ 
ô 
ˆŒõ ”I˜l¨1œo¨wÑ7Ô7ˆŒ	ˆ	ˆ	r-   r3   r[   c                 óª   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        S rq   )rg   r™  rš  r›  rj  s     r.   rp   z"Sapiens2PointmapFinalLayer.forward÷  sG   € ØŸš ]Ñ3Ô3ˆØŸš MÑ2Ô2ˆØŸš MÑ2Ô2ˆØ�yŠy˜Ñ'Ô'Ð'r-   )r   rZ  )r$   r%   r&   r    r(   r«   rG   r)   r   rp   rr   rs   s   @r.   r•  r•  ë  s�   ø€ € € € € ð	8ð 	8˜sð 	8°%¸¸S¸´/ð 	8ÈCð 	8Ðadð 	8ð 	8ð 	8ð 	8ð 	8ð 	8ð( U¤\ð (°e´lð (ð (ð (ð (ð (ð (ð (ð (r-   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 )ÚSapiens2PointmapScaleHeadrB   c                 óÎ  •— t          ¦   «                              ¦   «          t          j        ¦   «         | _        |j        g|j        j        d d…         z   }t          ||j        j        |j        j	        ¦  «        D ]8\  }}}| j         
                    t          |||d|dz
  dz  ¬¦  «        ¦  «         Œ9t          |j        j        |j        j        |j        ¬¦  «        | _        d S )Nr`   r_   r   )rD   rE   r]  )r_  )rF   rG   r   r{  r�  rJ   rq  Úscale_conv_out_channelsr|  Úscale_conv_kernel_sizesÚappendrY  r•  Úscale_final_input_sizeÚscale_final_hidden_sizesr4  rƒ  )rW   rB   Úscale_in_channelsru  rv  rD   rX   s         €r.   rG   z"Sapiens2PointmapScaleHead.__init__ÿ  sù   ø€ Ý‰Œ×ÒÑÔÐÝœ=™?œ?ˆÔØ#Ô/Ð0°6Ô3EÔ3]Ð^aÐ_aÐ^aÔ3bÑbÐÝ*-ØØÔÔ6ØÔÔ6ñ+
ô +
ð 	ð 	Ñ&ˆE�6˜;ð
 Ô×#Ò#Ý! %¨¸[ÐQRÐ]hÐklÑ]lÐqrÑ\rÐsÑsÔsñô ð ð õ 4ØÔÔ5ØÔÔ7ØÔ(ð
ñ 
ô 
ˆŒˆˆr-   r3   r[   c                 óV   — | j         D ]} ||¦  «        }Œ|                      |¦  «        S rq   )r�  rƒ  rˆ  s      r.   rp   z!Sapiens2PointmapScaleHead.forward  s7   € ØÔ%ð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMØ�~Š~˜mÑ,Ô,Ð,r-   rŠ  rs   s   @r.   rž  rž  þ  sj   ø€ € € € € ð
˜~ð 
ð 
ð 
ð 
ð 
ð 
ð$- U¤\ð -°e´lð -ð -ð -ð -ð -ð -ð -ð -r-   rž  c                   ó’   ‡ — e Zd ZU eed<   dZdZdZdZdgZ	dZ
dZdZdZeedœZdgZd	gZ ej        ¦   «         dˆ fd„¦   «         Zˆ xZS )ÚSapiens2PreTrainedModelrB   ÚmodelrY   )ÚimageTrJ  )r3   r4   ÚperiodsrN   r[   Nc                 ó  •— t          ¦   «                              |¦  «         t          |t          j        t          j        f¦  «        r(t          j        |j        d| j	        j
        ¬¦  «         dS t          |t          j        ¦  «        rt          j        |j        dd¬¦  «         dS t          |t          ¦  «        r…t          j        |j        d| j	        j
        ¬¦  «         |j	        j        dk    r&t          j        |j        d| j	        j
        ¬¦  «         |j	        j        rt          j        |j        ¦  «         dS dS t          |t(          ¦  «        r&t          j        |j        | j	        j        ¦  «         dS t          |t0          ¦  «        rQd|j        t5          j        ddd|j        z  t4          j        ¬	¦  «        z  z  }t          j        |j        |¦  «         dS t          |t@          tB          f¦  «        r„| "                    ¦   «         D ]q}t          |t          j        ¦  «        rt          j        |j        dd¬¦  «         Œ9t          |t          j        ¦  «        rt          j        |j        d
d¬¦  «         ŒpdS dS )zInitialize the weightsrÒ   )rÅ   ÚstdÚfan_outÚrelu)ÚmodeÚnonlinearityr   r   r–   r]   Úfan_inÚlinearN)#rF   Ú_init_weightsrŸ   r   r	  rR   ÚinitÚtrunc_normal_re   rB   Úinitializer_rangerd  Úkaiming_normal_rA   rK   rP   rQ   rL   Úzeros_rN   r%  Ú	constant_r(  r'  r“   r™   r)   r}   r›   rœ   Úcopy_r”   rm  rž  Úmodules)rW   rÓ   r”   Úhead_modulerX   s       €r.   r´  z%Sapiens2PreTrainedModel._init_weights-  sj  ø€ õ 	‰Œ×Ò˜fÑ%Ô%Ð%Ý�f�rœy­"¬)Ð4Ñ5Ô5ð 	cÝÔ˜vœ}°3¸D¼KÔ<YÐZÑZÔZÐZÐZÐZÝ˜¥Ô 2Ñ3Ô3ð 	cÝÔ  ¤°YÈVÐTÑTÔTÐTÐTÐTÝ˜Õ 2Ñ3Ô3ð 	cÝÔ˜vÔ/°c¸t¼{Ô?\Ð]Ñ]Ô]Ð]ØŒ}Ô0°1Ò4Ð4ÝÔ" 6Ô#9ÀÈÌÔIfÐgÑgÔgÐgØŒ}Ô+ð /Ý”˜FÔ-Ñ.Ô.Ð.Ð.Ð.ð/ð /å˜Õ 2Ñ3Ô3ð 
	cÝŒN˜6œ>¨4¬;Ô+GÑHÔHÐHÐHÐHÝ˜Õ =Ñ>Ô>ð 	cØ˜6œ;­%¬,°q¸!¸QÀÄÑ=PÕX]ÔXeÐ*fÑ*fÔ*fÑfÑfˆHÝŒJ�v”¨Ñ1Ô1Ð1Ð1Ð1Ý˜¥Õ/HÐ IÑJÔJð 	cØ%Ÿ~š~Ñ/Ô/ð cð c�Ý˜k­2¬9Ñ5Ô5ð cÝÔ(¨Ô);À)ÐZ`ÐaÑaÔaÐaÐaÝ ­R¬YÑ7Ô7ð cÝÔ(¨Ô);À(ÐYaÐbÑbÔbÐbøð	cð 	cðcð cr-   r+  )r$   r%   r&   r   r+   Úbase_model_prefixÚmain_input_nameÚinput_modalitiesÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_supports_sdpaÚ_supports_flash_attnÚ_supports_flex_attnÚ_supports_attention_backendrJ  rÿ   Ú_can_record_outputsÚ"_keys_to_ignore_on_load_unexpectedÚ_keys_to_ignore_on_load_missingr)   Úno_gradr´  rr   rs   s   @r.   r¨  r¨    sÃ   ø€ € € € € € àÐÐÑØÐØ$€OØ!ÐØ&*Ð#Ø(Ð)ÐØ€NØÐØÐØ"&Ðà&Ø'ðð Ðð +5¨Ð&à'4 oÐ#à€U„]�_„_ðcð cð cð cð cñ „_ðcð cð cð cð cr-   r¨  c                   ó´   ‡ — e Zd Zdefˆ fd„Ze ed¬¦  «        	 ddej        de	ej        ej        f         dz  de
e         d	efd
„¦   «         ¦   «         Zˆ xZS )ÚSapiens2EncoderrB   c                 óâ   •‡— t          ¦   «                              ‰¦  «         t          j        ˆfd„t	          ‰j        ¦  «        D ¦   «         ¦  «        | _        |                      ¦   «          d S )Nc                 ó2   •— g | ]}t          ‰|¬ ¦  «        ‘ŒS )rL  )rJ  )rt  r   rB   s     €r.   ú
<listcomp>z,Sapiens2Encoder.__init__.<locals>.<listcomp>L  s&   ø€ ÐiÐiÐi¸I�]˜6¨YÐ7Ñ7Ô7ÐiÐiÐir-   )rF   rG   r   r{  ÚrangeÚnum_hidden_layersr‰  Ú	post_initrV   s    `€r.   rG   zSapiens2Encoder.__init__I  si   øø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý”]ØiÐiÐiÐiÍÈvÔOgÑIhÔIhÐiÑiÔiñ
ô 
ˆŒ
ð 	�ŠÑÔÐÐÐr-   F)Útie_last_hidden_statesNr3   r  rå   r[   c                 óL   — | j         D ]} ||fd|i|¤Ž}Œt          |¬¦  «        S )Nr  )Úlast_hidden_state)r‰  r   )rW   r3   r  rå   Úlayer_modules        r.   rp   zSapiens2Encoder.forwardQ  sH   € ð !œJð 	kð 	kˆLØ(˜L¨ÐjÐjÐL_ÐjÐciÐjÐjˆMˆMå°Ð?Ñ?Ô?Ð?r-   rq   )r$   r%   r&   r   rG   r   r   r)   r   r(   r   r   r   rp   rr   rs   s   @r.   rÌ  rÌ  H  s×   ø€ € € € € ð˜~ð ð ð ð ð ð ð  Ø€_¨EÐ2Ñ2Ô2ð IMð	@ð 	@à”|ð	@ð # 5¤<°´Ð#=Ô>ÀÑEð	@ð Ð+Ô,ð		@ð
 
ð	@ð 	@ð 	@ñ 3Ô2ñ  Ôð	@ð 	@ð 	@ð 	@ð 	@r-   rÌ  c                   óŒ   ‡ — e Zd Zdefˆ fd„Zd„ Zee	 d
dej	        dej	        dz  de
e         defd	„¦   «         ¦   «         Zˆ xZS )ÚSapiens2ModelrB   c                 ó8  •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          |¦  «        | _        t          |¦  «        | _        t          |j	        |j
        ¬¦  «        | _        d| _        |                      ¦   «          d S )Nr  F)rF   rG   rA   ro   r“   Úrope_embeddingsrÌ  r©  r¼   rJ   r  rg  Úgradient_checkpointingrÒ  rV   s     €r.   rG   zSapiens2Model.__init__a  s�   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý,¨VÑ4Ô4ˆŒÝ<¸VÑDÔDˆÔÝ$ VÑ,Ô,ˆŒ
Ý# FÔ$6¸FÔ<OÐPÑPÔPˆŒ	Ø&+ˆÔ#à�ŠÑÔÐÐÐr-   c                 ó   — | j         j        S rq   ©ro   rU   rÊ   s    r.   Úget_input_embeddingsz"Sapiens2Model.get_input_embeddingsk  ó   € ØŒÔ/Ð/r-   NrY   rZ   rå   r[   c                 óV  — |                      | j        j        j        j        ¦  «        }|                      ||¬¦  «        }|                      |¦  «        } | j        ||fi |¤Ž}|                      |j        ¦  «        }|dd…ddd…f         }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). Only relevant for
            pre-training.

        Example:

        ```python
        >>> from transformers import AutoImageProcessor, AutoModel
        >>> from transformers.image_utils import load_image
        >>> import torch

        >>> image = load_image("http://images.cocodataset.org/val2017/000000004016.jpg")
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/sapiens2-pretrain-0.4b")
        >>> model = AutoModel.from_pretrained("facebook/sapiens2-pretrain-0.4b")

        >>> inputs = image_processor(images=image, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     outputs = model(**inputs)

        >>> cls_token = outputs.pooler_output
        >>> cls_token.shape
        torch.Size([1, 1024])
        ```
        )rZ   Nr   )rÕ  Úpooler_outputr3   r4   )rf   ro   rU   re   r^   rÚ  r©  rg  rÕ  r   r3   r4   )	rW   rY   rZ   rå   r3   r  ÚoutputÚsequence_outputÚpooled_outputs	            r.   rp   zSapiens2Model.forwardn  s»   € ðD $—’ t¤Ô'GÔ'NÔ'TÑUÔUˆØŸš¨Ào˜ÑVÔVˆØ"×2Ò2°<Ñ@Ô@Ðà�”˜MÐ+>ÐIÐIÀ&ÐIÐIˆØŸ)š) FÔ$<Ñ=Ô=ˆØ'¨¨¨¨1¨a¨a¨a¨Ô0ˆå)Ø-Ø'Ø Ô.ØÔ(ð	
ñ 
ô 
ð 	
r-   rq   )r$   r%   r&   r   rG   rÞ  r   r   r)   r   r   r   r   rp   rr   rs   s   @r.   rØ  rØ  _  s½   ø€ € € € € ð˜~ð ð ð ð ð ð ð0ð 0ð 0ð Øð 04ð-
ð -
à”lð-
ð œ¨Ñ,ð-
ð Ð+Ô,ð	-
ð
 
$ð-
ð -
ð -
ñ „^ñ Ôð-
ð -
ð -
ð -
ð -
r-   rØ  c            	       ó„   ‡ — e Zd Zdefˆ fd„Zd„ Zeeede	j
        dee         defd„¦   «         ¦   «         ¦   «         Zˆ xZS )ÚSapiens2BackbonerB   c                 óŠ  •‡— t          ¦   «                              ‰¦  «         t          ‰¦  «        | _        t	          ‰¦  «        | _        t          ‰¦  «        | _        t          ‰j	        ‰j
        ¬¦  «        | _        d| _        ˆfd„t          ‰j        dz   ¦  «        D ¦   «         | _        |                      ¦   «          d S )Nr  Fc                 ó   •— g | ]	}‰j         ‘Œ
S r,   )rJ   )rt  rµ   rB   s     €r.   rÏ  z-Sapiens2Backbone.__init__.<locals>.<listcomp>«  s   ø€ Ð]Ð]Ð]°A˜VÔ/Ð]Ð]Ð]r-   r   )rF   rG   rA   ro   r“   rÚ  rÌ  r©  r¼   rJ   r  rg  rÛ  rÐ  rÑ  Únum_featuresrÒ  rV   s    `€r.   rG   zSapiens2Backbone.__init__¢  s¯   øø€ Ý‰Œ×Ò˜Ñ Ô Ð å,¨VÑ4Ô4ˆŒÝ<¸VÑDÔDˆÔÝ$ VÑ,Ô,ˆŒ
Ý# FÔ$6¸FÔ<OÐPÑPÔPˆŒ	Ø&+ˆÔ#à]Ð]Ð]Ð]½¸vÔ?WÐZ[Ñ?[Ñ9\Ô9\Ð]Ñ]Ô]ˆÔØ�ŠÑÔÐÐÐr-   c                 ó   — | j         j        S rq   rÝ  rÊ   s    r.   rÞ  z%Sapiens2Backbone.get_input_embeddings®  rß  r-   rY   rå   r[   c                 ól  — |                      | j        j        j        j        ¦  «        }|                      |¦  «        }|                      |¦  «        }d|d<    | j        ||fi |¤Ž}|j        }|j        \  }}}	}
| j	        j
        }t          |t          ¦  «        r|n|d         }t          |t          ¦  «        r|n|d         }|	|z  }|
|z  }dt          | j	        dd¦  «        z   }t          | j	        dd¦  «        }g g }}t          t          | j        |¦  «        ¦  «        D ]Ö\  }\  }}| j	        j        r|                      |¦  «        }|| j        v r¤|r"|                     |dd…ddd…f         ¦  «         |dd…|d…dd…f         }| j	        j        rL|                     ||||j        d	         ¦  «                             dd
dd¦  «                             ¦   «         }n|}|                     |¦  «         Œ×t3          t5          |¦  «        |rt5          |¦  «        nd|j        |j        ¬¦  «        S )a2  
        Example:

        ```python
        >>> from transformers import AutoBackbone, AutoImageProcessor
        >>> from transformers.image_utils import load_image
        >>> import torch

        >>> image = load_image("http://images.cocodataset.org/val2017/000000004016.jpg")
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/sapiens2-pretrain-0.4b")
        >>> model = AutoBackbone.from_pretrained("facebook/sapiens2-pretrain-0.4b")

        >>> inputs = image_processor(images=image, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     outputs = model(**inputs, return_class_token=True)

        >>> outputs.feature_maps[0].shape
        torch.Size([1, 1024, 64, 48])
        >>> outputs.cls_tokens[0].shape
        torch.Size([1, 1024])
        ```
        TÚoutput_hidden_statesr   r   rP   Úreturn_class_tokenFNr`   r   r_   )Úfeature_mapsr#   r3   r4   )rf   ro   rU   re   r^   rÚ  r©  r3   rd   rB   rT   rŸ   r    ÚgetattrÚ	enumerater|  Ústage_namesÚnormalize_backbone_outputsrg  Úout_featuresr¢  Úreshape_hidden_statesrú   Úpermuterä   r"   r(   r4   )rW   rY   rå   r3   r  râ  Ústage_hidden_statesrm   rµ   Úimage_heightÚimage_widthrT   r£   r¤   Únum_patches_heightÚnum_patches_widthÚ
num_prefixrí  rî  r#   ÚidxÚ
stage_namer)  Úpatch_tokensÚfeature_maps                            r.   rp   zSapiens2Backbone.forward±  sp  € ð< $—’ t¤Ô'GÔ'NÔ'TÑUÔUˆØŸš¨Ñ5Ô5ˆØ"×2Ò2°<Ñ@Ô@Ðà)-ˆÐ%Ñ&Ø�”˜MÐ+>ÐIÐIÀ&ÐIÐIˆØ$Ô2Ðà3?Ô3EÑ0ˆ
�A�| [Ø”[Ô+ˆ
Ý%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆÝ%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆØ)¨\Ñ9ÐØ'¨<Ñ7Ðà� ¤Ð.CÀQÑGÔGÑGˆ
Ý$ T¤[Ð2FÈÑNÔNÐà#% r�jˆÝ/8½¸TÔ=MÐObÑ9cÔ9cÑ/dÔ/dð 	1ð 	1Ñ+ˆCÑ+�*˜lØŒ{Ô5ð 7Ø#Ÿyšy¨Ñ6Ô6�à˜TÔ.Ð.Ð.Ø%ð =Ø×%Ò% l°1°1°1°a¸¸¸°7Ô&;Ñ<Ô<Ð<Ø+¨A¨A¨A¨z¨{¨{¸A¸A¸AÐ,=Ô>�Ø”;Ô4ð /à$×,Ò,¨ZÐ9KÐM^Ð`lÔ`rÐsuÔ`vÑwÔwß š  A q¨!Ñ,Ô,ß#š™œð  �Kð #/�Kà×#Ò# KÑ0Ô0Ð0øå%Ý˜|Ñ,Ô,Ø,>ÐH•u˜ZÑ(Ô(Ð(ÀDØ Ô.ØÔ(ð	
ñ 
ô 
ð 	
r-   )r$   r%   r&   r   rG   rÞ  r   r   r   r)   r   r   r   r"   rp   rr   rs   s   @r.   ræ  ræ     s¸   ø€ € € € € ð
˜~ð 
ð 
ð 
ð 
ð 
ð 
ð0ð 0ð 0ð Ø ØðF
à”lðF
ð Ð+Ô,ðF
ð 
 ð	F
ð F
ð F
ñ „^ñ !Ô ñ ÔðF
ð F
ð F
ð F
ð F
r-   ræ  zfacebook/sapiens2-seg-0.4b)Ú
checkpointc                   ó†   ‡ — e Zd Zdefˆ fd„Zee	 d	dej        dej	        dz  de
e         defd„¦   «         ¦   «         Zˆ xZS )
ÚSapiens2ForSemanticSegmentationrB   c                 óÚ   •— t          ¦   «                              |¦  «         |j        | _        t          |¦  «        | _        t          |¦  «        | _        |                      ¦   «          d S rq   ©rF   rG   r‚  rØ  r©  rm  Údecode_headrÒ  rV   s     €r.   rG   z(Sapiens2ForSemanticSegmentation.__init__ÿ  óZ   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø Ô+ˆŒÝ" 6Ñ*Ô*ˆŒ
Ý'¨Ñ/Ô/ˆÔØ�ŠÑÔÐÐÐr-   NrY   Úlabelsrå   r[   c                 óh  — |�| j         j        dk    rt          d¦  «        ‚ | j        |fi |¤Ž}|j        \  }}}}| j         j        }	t          |	t          ¦  «        r|	n|	d         }
t          |	t          ¦  «        r|	n|	d         }||
z  }||z  }|j        dd…d| j         j	        z   d…f         }| 
                    dd¦  «                             |d||¦  «        }|                      |¦  «        }d}|�"|                      ||| j         j        ¬¦  «        }t          |||j        |j        ¬¦  «        S )	aø  
        labels (`torch.LongTensor` of shape `(batch_size, height, width)`, *optional*):
            Ground truth semantic segmentation maps for computing the loss.
            Indices should be in `[0, ..., config.num_labels - 1]`.
            If `config.num_labels > 1`, a classification loss is computed (Cross-Entropy).

        Example:

        ```python
        >>> from transformers import AutoImageProcessor, AutoModel
        >>> from transformers.image_utils import load_image
        >>> import torch

        >>> image = load_image("http://images.cocodataset.org/val2017/000000004016.jpg")
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/sapiens2-seg-0.4b")
        >>> model = AutoModel.from_pretrained("facebook/sapiens2-seg-0.4b")

        >>> inputs = image_processor(image, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     outputs = model(**inputs)

        >>> outputs.logits.shape
        torch.Size([1, 29, 1024, 768])
        ```
        Nr   z/The number of labels should be greater than oner   r_   r`   )Úignore_index)r1   Úlogitsr3   r4   )rB   r‚  rc   r©  rd   rT   rŸ   r    rÕ  rP   rh   rú   r  Úloss_functionÚsemantic_loss_ignore_indexr   r3   r4   )rW   rY   r  rå   Úoutputsrm   rµ   r¶   r·   rT   r£   r¤   Úpatch_heightÚpatch_widthrþ  rÿ  r
  r1   s                     r.   rp   z'Sapiens2ForSemanticSegmentation.forward  sa  € ðB Ð $¤+Ô"8¸AÒ"=Ð"=ÝÐNÑOÔOÐOà�$”*˜\Ð4Ð4¨VÐ4Ð4ˆà'3Ô'9Ñ$ˆ
�A�v˜uØ”[Ô+ˆ
Ý%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆÝ%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆØ Ñ-ˆØ˜|Ñ+ˆàÔ0°°°°A¸¼Ô8WÑ4WÐ4YÐ4YÐ1YÔZˆØ"×,Ò,¨Q°Ñ2Ô2×:Ò:¸:ÀrÈ<ÐYdÑeÔeˆà×!Ò! +Ñ.Ô.ˆàˆØÐØ×%Ò% f¨fÀ4Ä;ÔCiÐ%ÑjÔjˆDå&ØØØ!Ô/ØÔ)ð	
ñ 
ô 
ð 	
r-   rq   )r$   r%   r&   r   rG   r   r   r)   r*   Ú
LongTensorr   r   r   rp   rr   rs   s   @r.   r  r  ý  s°   ø€ € € € € ð˜~ð ð ð ð ð ð ð Øð +/ð9
ð 9
àÔ'ð9
ð Ô  4Ñ'ð9
ð Ð+Ô,ð	9
ð
 
!ð9
ð 9
ð 9
ñ „^ñ Ôð9
ð 9
ð 9
ð 9
ð 9
r-   r  úgaussian-heatmapc                 ó&  — |dvrt          d¦  «        ‚| j        dk    rt          d¦  «        ‚| j        \  }}}}d}|dk    r2d}|                      ¦   «         } | dd…ddd…d	f          | dd…ddd…d	f<   |                      |d
|||¦  «        } |                      ¦   «         }|                     d
¦  «        \  }	}
| dd…|
d	f         |dd…|	d	f<   | dd…|	d	f         |dd…|
d	f<   |                     ||||f¦  «        }|                     d
¦  «        }|S )aÃ  Flip the flipped heatmaps back to the original form.

    Args:
        output_flipped (`torch.tensor` of shape `(batch_size, num_keypoints, height, width)`):
            The output heatmaps obtained from the flipped images.
        flip_pairs (`torch.Tensor` of shape `(num_keypoints, 2)`):
            Pairs of keypoints which are mirrored (for example, left ear -- right ear).
        target_type (`str`, *optional*, defaults to `"gaussian-heatmap"`):
            Target type to use. Can be gaussian-heatmap or combined-target.
            gaussian-heatmap: Classification target with gaussian distribution.
            combined-target: The combination of classification target (response map) and regression target (offset map).
            Paper ref: Huang et al. The Devil is in the Details: Delving into Unbiased Data Processing for Human Pose Estimation (CVPR 2020).

    Returns:
        torch.Tensor: heatmaps that flipped back to the original image
    )r  úcombined-targetz9target_type should be gaussian-heatmap or combined-targetr–   zCoutput_flipped should be [batch_size, num_keypoints, height, width]r   r  r   N.r`   )rc   rB  rd   Úclonerú   ÚunbindÚflip)Úoutput_flippedÚ
flip_pairsÚtarget_typerm   Únum_keypointsr¶   r·   ÚchannelsÚoutput_flipped_backÚleft_indicesÚright_indicess              r.   Ú	flip_backr  D  st  € ð" ÐAÐAÐAÝÐTÑUÔUÐUàÔ˜aÒÐÝÐ^Ñ_Ô_Ð_à/=Ô/CÑ,€J�˜v uØ€HØÐ'Ò'Ð'ØˆØ'×-Ò-Ñ/Ô/ˆØ(6°q°q°q¸!¸$¸Q¸$À°|Ô(DÐ'Dˆ�q�q�q˜!˜$˜Q˜$ �|Ñ$Ø#×+Ò+¨J¸¸HÀfÈeÑTÔT€NØ(×.Ò.Ñ0Ô0Ðð #-×"3Ò"3°BÑ"7Ô"7Ñ€L�-Ø0>¸q¸q¸qÀ-ÐQTÐ?TÔ0UÐ˜˜˜˜<¨Ð,Ñ-Ø1?ÀÀÀÀ<ÐQTÐ@TÔ1UÐ˜˜˜˜=¨#Ð-Ñ.Ø-×5Ò5°zÀ=ÐRXÐZ_Ð6`ÑaÔaÐà-×2Ò2°2Ñ6Ô6ÐØÐr-   zfacebook/sapiens2-pose-0.4bz�
    The Sapiens2 model with a pose estimation head on top (a set of heatmap predictors on top of the hidden states output).
    )r   r    c                   ó²   ‡ — e Zd Zdefˆ fd„Zee	 	 	 ddej        dej	        dz  dej        dz  dej        dz  de
e         d	efd
„¦   «         ¦   «         Zˆ xZS )ÚSapiens2ForPoseEstimationrB   c                 óÚ   •— t          ¦   «                              |¦  «         |j        | _        t          |¦  «        | _        t          |¦  «        | _        |                      ¦   «          d S rq   r  rV   s     €r.   rG   z"Sapiens2ForPoseEstimation.__init__u  r  r-   NrY   r  r  Úlabel_weightsrå   r[   c                 ó4  —  | j         |fi |¤Ž}|j        \  }}}	}
| j        j        }t	          |t
          ¦  «        r|n|d         }t	          |t
          ¦  «        r|n|d         }|	|z  }|
|z  }|j        dd…d| j        j        z   d…f         }|                     dd¦  «         	                    |d||¦  «        }|  
                    |¦  «        }|�t          ||¦  «        }d}|�t          j        |||¬¦  «        }t          |||j        |j        ¬¦  «        S )a   
        flip_pairs (`torch.Tensor` of shape `(num_pairs, 2)`, *optional*):
            Pairs of keypoints which are mirrored (for example, left ear -- right ear), used for
            test-time flip augmentation. When provided, the model assumes `pixel_values` contains
            horizontally-flipped images and calls `flip_back` on the output heatmaps to restore the
            original orientation.
        labels (`torch.FloatTensor` of shape `(batch_size, num_keypoints, height, width)`, *optional*):
            Heatmap ground truth for computing the loss.
        label_weights (`torch.FloatTensor` of shape `(batch_size, num_labels, 1, 1)` or `(batch_size, num_labels, height, width)`, *optional*):
            Visibility weights for each keypoint. Must be broadcastable to the shape of `labels`.

        Example:

        ```python
        >>> from transformers import AutoImageProcessor, AutoModel
        >>> from transformers.image_utils import load_image
        >>> import torch

        >>> image = load_image("http://images.cocodataset.org/val2017/000000004016.jpg")
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/sapiens2-pose-0.4b")
        >>> model = AutoModel.from_pretrained("facebook/sapiens2-pose-0.4b")

        >>> boxes = [[[270.8, 0.6, 294.1, 379.5]]]
        >>> inputs = image_processor(image, boxes=boxes, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     outputs = model(**inputs)

        >>> outputs.heatmaps.shape
        torch.Size([1, 308, 256, 192])
        ```
        r   r   Nr_   r`   )re   )r1   r2   r3   r4   )r©  rd   rB   rT   rŸ   r    rÕ  rP   rh   rú   r  r  ÚFÚmse_lossr0   r3   r4   )rW   rY   r  r  r#  rå   r  rm   rµ   r¶   r·   rT   r£   r¤   r  r  rþ  rÿ  r2   r1   s                       r.   rp   z!Sapiens2ForPoseEstimation.forward|  sI  € ðR �$”*˜\Ð4Ð4¨VÐ4Ð4ˆà'3Ô'9Ñ$ˆ
�A�v˜uØ”[Ô+ˆ
Ý%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆÝ%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆØ Ñ-ˆØ˜|Ñ+ˆàÔ0°°°°A¸¼Ô8WÑ4WÐ4YÐ4YÐ1YÔZˆØ"×,Ò,¨Q°Ñ2Ô2×:Ò:¸:ÀrÈ<ÐYdÑeÔeˆà×#Ò# KÑ0Ô0ˆØÐ!Ý  ¨:Ñ6Ô6ˆHàˆØÐÝ”:˜h¨°}ÐEÑEÔEˆDå*ØØØ!Ô/ØÔ)ð	
ñ 
ô 
ð 	
r-   ©NNN)r$   r%   r&   r   rG   r   r   r)   r*   r   r   r   r0   rp   rr   rs   s   @r.   r!  r!  n  sè   ø€ € € € € ð˜~ð ð ð ð ð ð ð Øð +/Ø+/Ø26ð@
ð @
àÔ'ð@
ð ”L 4Ñ'ð@
ð Ô! DÑ(ð	@
ð
 Ô(¨4Ñ/ð@
ð Ð+Ô,ð@
ð 
%ð@
ð @
ð @
ñ „^ñ Ôð@
ð @
ð @
ð @
ð @
r-   r!  zfacebook/sapiens2-normal-0.4bzƒ
    The Sapiens2 model with a normal estimation head on top (a PixelShuffle-based decoder that predicts surface normal maps).
    c                   ó†   ‡ — e Zd Zdefˆ fd„Zee	 d	dej        dej        dz  de	e
         defd„¦   «         ¦   «         Zˆ xZS )
ÚSapiens2ForNormalEstimationrB   c                 óÚ   •— t          ¦   «                              |¦  «         |j        | _        t          |¦  «        | _        t          |¦  «        | _        |                      ¦   «          d S rq   r  rV   s     €r.   rG   z$Sapiens2ForNormalEstimation.__init__È  r  r-   NrY   r  rå   r[   c                 ó   —  | j         |fi |¤Ž}|j        \  }}}}| j        j        }	t	          |	t
          ¦  «        r|	n|	d         }
t	          |	t
          ¦  «        r|	n|	d         }||
z  }||z  }|j        dd…d| j        j        z   d…f         }|                     dd¦  «         	                    |d||¦  «        }|  
                    |¦  «        }d}|�t          d¦  «        ‚t          |||j        |j        ¬¦  «        S )ae  
        labels (`torch.FloatTensor` of shape `(batch_size, num_labels, height, width)`, *optional*):
            Ground-truth surface normal maps for computing the loss.

        Example:

        ```python
        >>> from transformers import AutoImageProcessor, AutoModel
        >>> from transformers.image_utils import load_image
        >>> import torch

        >>> image = load_image("http://images.cocodataset.org/val2017/000000004016.jpg")
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/sapiens2-normal-0.4b")
        >>> model = AutoModel.from_pretrained("facebook/sapiens2-normal-0.4b")

        >>> inputs = image_processor(image, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     outputs = model(**inputs)

        >>> outputs.normals.shape
        torch.Size([1, 3, 1024, 768])
        ```
        r   r   Nr_   r`   úTraining is not yet supported)r1   r7   r3   r4   )r©  rd   rB   rT   rŸ   r    rÕ  rP   rh   rú   r  ÚNotImplementedErrorr6   r3   r4   )rW   rY   r  rå   r  rm   rµ   r¶   r·   rT   r£   r¤   r  r  rþ  rÿ  r7   r1   s                     r.   rp   z#Sapiens2ForNormalEstimation.forwardÏ  s,  € ð> �$”*˜\Ð4Ð4¨VÐ4Ð4ˆà'3Ô'9Ñ$ˆ
�A�v˜uØ”[Ô+ˆ
Ý%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆÝ%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆØ Ñ-ˆØ˜|Ñ+ˆàÔ0°°°°A¸¼Ô8WÑ4WÐ4YÐ4YÐ1YÔZˆØ"×,Ò,¨Q°Ñ2Ô2×:Ò:¸:ÀrÈ<ÐYdÑeÔeˆà×"Ò" ;Ñ/Ô/ˆàˆØÐÝ%Ð&EÑFÔFÐFå,ØØØ!Ô/ØÔ)ð	
ñ 
ô 
ð 	
r-   rq   )r$   r%   r&   r   rG   r   r   r)   r*   r   r   r6   rp   rr   rs   s   @r.   r)  r)  Á  s°   ø€ € € € € ð˜~ð ð ð ð ð ð ð Øð ,0ð4
ð 4
àÔ'ð4
ð Ô! DÑ(ð4
ð Ð+Ô,ð	4
ð
 
'ð4
ð 4
ð 4
ñ „^ñ Ôð4
ð 4
ð 4
ð 4
ð 4
r-   r)  zfacebook/sapiens2-pointmap-0.4bzÅ
    The Sapiens2 model with a pointmap head on top (a PixelShuffle-based decoder that predicts per-pixel 3D XYZ
    coordinates, plus an optional scale branch for focal-length normalization).
    c                   ó†   ‡ — e Zd Zdefˆ fd„Zee	 d	dej        dej        dz  de	e
         defd„¦   «         ¦   «         Zˆ xZS )
ÚSapiens2ForPointmapEstimationrB   c                 ó6  •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          |¦  «        | _        |j        �|j        j        �t          |¦  «        nt          j
        ¦   «         | _        |                      ¦   «          d S rq   )rF   rG   rØ  r©  rm  r  rq  r   rž  r   r  Ú
scale_headrÒ  rV   s     €r.   rG   z&Sapiens2ForPointmapEstimation.__init__  sˆ   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý" 6Ñ*Ô*ˆŒ
Ý'¨Ñ/Ô/ˆÔð Ô!Ð-°&Ô2DÔ2\Ð2hõ & fÑ-Ô-Ð-å”‘”ð 	Œð
 	�ŠÑÔÐÐÐr-   NrY   r  rå   r[   c                 ón  —  | j         |fi |¤Ž}|j        \  }}}}| j        j        }	t	          |	t
          ¦  «        r|	n|	d         }
t	          |	t
          ¦  «        r|	n|	d         }||
z  }||z  }|j        dd…d| j        j        z   d…f         }|                     dd¦  «         	                    |d||¦  «        }|  
                    |¦  «        }t	          | j        t          j        ¦  «        rdn|                      |¦  «        }d}|�t          d¦  «        ‚t          ||||j        |j        ¬¦  «        S )aW  
        labels (`torch.FloatTensor` of shape `(batch_size, 3, height, width)`, *optional*):
            Ground-truth pointmap for computing the loss.

        Example:

        ```python
        >>> from transformers import AutoImageProcessor, AutoModel
        >>> from transformers.image_utils import load_image
        >>> import torch

        >>> image = load_image("http://images.cocodataset.org/val2017/000000004016.jpg")
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/sapiens2-pointmap-0.4b")
        >>> model = AutoModel.from_pretrained("facebook/sapiens2-pointmap-0.4b")

        >>> inputs = image_processor(image, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     outputs = model(**inputs)

        >>> outputs.pointmaps.shape
        torch.Size([1, 3, 1024, 768])
        ```
        r   r   Nr_   r`   r,  )r1   r:   r;   r3   r4   )r©  rd   rB   rT   rŸ   r    rÕ  rP   rh   rú   r  r1  r   r  r-  r9   r3   r4   )rW   rY   r  rå   r  rm   rµ   r¶   r·   rT   r£   r¤   r  r  rþ  rÿ  r:   r;   r1   s                      r.   rp   z%Sapiens2ForPointmapEstimation.forward  sX  € ð> �$”*˜\Ð4Ð4¨VÐ4Ð4ˆà'3Ô'9Ñ$ˆ
�A�v˜uØ”[Ô+ˆ
Ý%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆÝ%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆØ Ñ-ˆØ˜|Ñ+ˆàÔ0°°°°A¸¼Ô8WÑ4WÐ4YÐ4YÐ1YÔZˆØ"×,Ò,¨Q°Ñ2Ô2×:Ò:¸:ÀrÈ<ÐYdÑeÔeˆà×$Ò$ [Ñ1Ô1ˆ	Ý# D¤OµR´[ÑAÔAÐc��ÀtÇÂÐWbÑGcÔGcˆàˆØÐÝ%Ð&EÑFÔFÐFå.ØØØØ!Ô/ØÔ)ð
ñ 
ô 
ð 	
r-   rq   )r$   r%   r&   r   rG   r   r   r)   r*   r   r   r9   rp   rr   rs   s   @r.   r/  r/    s°   ø€ € € € € ð	˜~ð 	ð 	ð 	ð 	ð 	ð 	ð Øð ,0ð6
ð 6
àÔ'ð6
ð Ô! DÑ(ð6
ð Ð+Ô,ð	6
ð
 
)ð6
ð 6
ð 6
ñ „^ñ Ôð6
ð 6
ð 6
ð 6
ð 6
r-   r/  zfacebook/sapiens2-matting-1bzœ
    The Sapiens2 model with a matting head on top (a PixelShuffle-based decoder that predicts a
    pre-multiplied RGB foreground and an alpha matte).
    c                   ó†   ‡ — e Zd Zdefˆ fd„Zee	 d	dej        dej        dz  de	e
         defd„¦   «         ¦   «         Zˆ xZS )
ÚSapiens2ForImageMattingrB   c                 óÂ   •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          |¦  «        | _        |                      ¦   «          d S rq   )rF   rG   rØ  r©  rm  r  rÒ  rV   s     €r.   rG   z Sapiens2ForImageMatting.__init__^  sP   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý" 6Ñ*Ô*ˆŒ
Ý'¨Ñ/Ô/ˆÔØ�ŠÑÔÐÐÐr-   NrY   r  rå   r[   c                 ó^  —  | j         |fi |¤Ž}|j        \  }}}}| j        j        }	t	          |	t
          ¦  «        r|	n|	d         }
t	          |	t
          ¦  «        r|	n|	d         }||
z  }||z  }|j        dd…d| j        j        z   d…f         }|                     dd¦  «         	                    |d||¦  «        }|  
                    |¦  «                             ¦   «         }|dd…dd…f         }|dd…dd…f         }d}|�t          d¦  «        ‚t          ||||j        |j        ¬¦  «        S )	a¡  
        labels (`torch.FloatTensor` of shape `(batch_size, 4, height, width)`, *optional*):
            Ground-truth matting targets for computing the loss.

        Example:

        ```python
        >>> from transformers import AutoImageProcessor, AutoModel
        >>> from transformers.image_utils import load_image
        >>> import torch

        >>> image = load_image("http://images.cocodataset.org/val2017/000000004016.jpg")
        >>> image_processor = AutoImageProcessor.from_pretrained("facebook/sapiens2-matting-1b")
        >>> model = AutoModel.from_pretrained("facebook/sapiens2-matting-1b")

        >>> inputs = image_processor(image, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     outputs = model(**inputs)

        >>> outputs.alphas.shape
        torch.Size([1, 1, 1024, 768])
        >>> outputs.foregrounds.shape
        torch.Size([1, 3, 1024, 768])
        ```
        r   r   Nr_   r`   r   r,  )r1   r>   r?   r3   r4   )r©  rd   rB   rT   rŸ   r    rÕ  rP   rh   rú   r  Úsigmoidr-  r=   r3   r4   )rW   rY   r  rå   r  rm   rµ   r¶   r·   rT   r£   r¤   r  r  rþ  rÿ  Úmattingr?   r>   r1   s                       r.   rp   zSapiens2ForImageMatting.forwardd  sf  € ðB �$”*˜\Ð4Ð4¨VÐ4Ð4ˆà'3Ô'9Ñ$ˆ
�A�v˜uØ”[Ô+ˆ
Ý%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆÝ%/°
½CÑ%@Ô%@ÐS�z�zÀjÐQRÄmˆØ Ñ-ˆØ˜|Ñ+ˆàÔ0°°°°A¸¼Ô8WÑ4WÐ4YÐ4YÐ1YÔZˆØ"×,Ò,¨Q°Ñ2Ô2×:Ò:¸:ÀrÈ<ÐYdÑeÔeˆà×"Ò" ;Ñ/Ô/×7Ò7Ñ9Ô9ˆØ˜a˜a˜a  ! ˜e”nˆØ˜˜˜˜A˜B˜B˜”ˆàˆØÐÝ%Ð&EÑFÔFÐFå)ØØØ#Ø!Ô/ØÔ)ð
ñ 
ô 
ð 	
r-   rq   )r$   r%   r&   r   rG   r   r   r)   r*   r   r   r=   rp   rr   rs   s   @r.   r4  r4  V  s°   ø€ € € € € ð˜~ð ð ð ð ð ð ð Øð ,0ð9
ð 9
àÔ'ð9
ð Ô! DÑ(ð9
ð Ð+Ô,ð	9
ð
 
$ð9
ð 9
ð 9
ñ „^ñ Ôð9
ð 9
ð 9
ð 9
ð 9
r-   r4  )r  r!  r)  r/  r4  rØ  r¨  ræ  r'  )rÒ   NN)r  )\r°   Úcollections.abcr   r   Údataclassesr   Únumpyr‰   r)   Útorch.nn.functionalr   râ   r%  r   Ú r	   rµ  Úactivationsr
   Úbackbone_utilsr   r   Úintegrationsr   Úmodeling_layersr   Úmodeling_outputsr   r   r   r   r   Úmodeling_utilsr   r   Úprocessing_utilsr   Úpytorch_utilsr   Úutilsr   r   Úutils.genericr   r   r   Úutils.output_capturingr   Úconfiguration_sapiens2r   r"   r0   r6   r9   r=   r“  rA   r    r^   rx   rƒ   rÌ   r‘   r“   r¼   rÑ   r(   rê   r÷   rÞ   rÿ   r%  r-  r9  r>  rJ  rY  rm  rŒ  r•  rž  r¨  rÌ  rØ  ræ  r  r  r!  r)  r/  r4  Ú__all__r,   r-   r.   ú<module>rK     sã
  ðð& €€€Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø !Ð !Ð !Ð !Ð !Ð !à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø HÐ HÐ HÐ HÐ HÐ HÐ HÐ HØ 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9ðð ð ð ð ð ð ð ð ð ð ð ð ð ð GÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø @Ð @Ð @Ð @Ð @Ð @Ø 7Ð 7Ð 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø YÐ YÐ YÐ YÐ YÐ YÐ YÐ YÐ YÐ YØ 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø 2Ð 2Ð 2Ð 2Ð 2Ð 2ð €ððñ ô ð ð7ð 7ð 7ð 7ð 7˜^ñ 7ô 7ñ „ñô ð7ð €ððñ ô ð
 ð<ð <ð <ð <ð < +ñ <ô <ñ „ñô ð<ð$ €ððñ ô ð
 ð<ð <ð <ð <ð < Kñ <ô <ñ „ñô ð<ð* €ððñ ô ð
 ð<ð <ð <ð <ð < kñ <ô <ñ „ñô ð<ð0 €ððñ ô ð
 ð1ð 1ð 1ð 1ð 1 ñ 1ô 1ñ „ñô ð1ð,"ð "ð "ð "ð "˜œñ "ô "ð "ðJ %Ð$¨RÐ0Ñ0Ô0ðØðØ'*ðØ38´;ðØHMÌðà
„\ðð ð ñ 1Ô0ðð< ØØ ð	ð ØŒLðà�4‰<ðð �D‰Lðð �T‰\ð	ð
 „\ðð ð ð ð:48ð 48ð 48ð 48ð 48 B¤Iñ 48ô 48ð 48ðn Ð˜YÑ'Ô'ðJð Jð Jð Jð J�b”iñ Jô Jñ (Ô'ðJð((ð (ð (ð Ø Ø ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð �S‰[ð%ð �T‰\ð%ð �T‰\ð%ð ˆ5Œ<˜œÐ%Ô&ð%ð %ð %ð %ðDØ„|ðØœðØ+0¬<ðØ>C¼lðà
ˆ5Œ<˜œÐ%Ô&ðð ð ð ðB	U˜Uœ\ð 	U°#ð 	U¸%¼,ð 	Uð 	Uð 	Uð 	Uð?)ð ?)ð ?)ð ?)ð ?)˜œ	ñ ?)ô ?)ð ?)ðD+ð +ð +ð +ð +˜œñ +ô +ð +ð<ð <ð <ð <ð <�"”)ñ <ô <ð <ðð ð ð ð �r”yñ ô ð ð %ð %ð %ð %ð %�r”yñ %ô %ð %ð0+ð +ð +ð +ð +Ð.ñ +ô +ð +ð\Bð Bð Bð Bð B˜œ	ñ Bô Bð BðJ5-ð 5-ð 5-ð 5-ð 5-�2”9ñ 5-ô 5-ð 5-ðp	ð 	ð 	ð 	ð 	 b¤iñ 	ô 	ð 	ð(ð (ð (ð (ð ( ¤ñ (ô (ð (ð&-ð -ð -ð -ð - ¤	ñ -ô -ð -ð2 ð-cð -cð -cð -cð -c˜oñ -cô -cñ „ð-cð`@ð @ð @ð @ð @Ð-ñ @ô @ð @ð. ð=
ð =
ð =
ð =
ð =
Ð+ñ =
ô =
ñ „ð=
ð@ ðY
ð Y
ð Y
ð Y
ð Y
�}Ð&=ñ Y
ô Y
ñ „ðY
ðx €Ð7Ð8Ñ8Ô8ðC
ð C
ð C
ð C
ð C
Ð&=ñ C
ô C
ñ 9Ô8ðC
ðL'ð 'ð 'ð 'ðT €Ø,ððñ ô ðJ
ð J
ð J
ð J
ð J
Ð 7ñ J
ô J
ñô ðJ
ðZ €Ø.ððñ ô ð>
ð >
ð >
ð >
ð >
Ð"9ñ >
ô >
ñô ð>
ðB €Ø0ððñ ô ðD
ð D
ð D
ð D
ð D
Ð$;ñ D
ô D
ñô ðD
ðN €Ø-ððñ ô ðB
ð B
ð B
ð B
ð B
Ð5ñ B
ô B
ñô ðB
ðJ	ð 	ð 	€€€r-   