§
    ‚ŠtjÊ ã                   óœ  — 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	Z	 d dl
m	c mZ d dlmZ ddl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 ddlmZ ddlmZ ddlm Z m!Z!m"Z"m#Z# ddl$m%Z%m&Z&m'Z' ddl(m)Z)m*Z* ddl+m,Z, ddl-m.Z.m/Z/m0Z0m1Z1m2Z2  e#j3        e4¦  «        Z5 e!d¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z6 e!d¬¦  «        e G d„ de ¦  «        ¦   «         ¦   «         Z7 G d„ de	j8        ¦  «        Z9 G d„ de	j8        ¦  «        Z: G d„ d e	j8        ¦  «        Z;	 dTd"e	j8        d#ej        d$ej        d%ej        d&ej        dz  d'e<d(e<fd)„Z=dUd*ej        d+e>dz  d,ej        fd-„Z? G d.„ d/e	j8        ¦  «        Z@ G d0„ d1e	j8        ¦  «        ZAd2„ ZBd3„ ZC G d4„ d5e¦  «        ZD e!d6¬¦  «        e G d7„ d8e ¦  «        ¦   «         ¦   «         ZEe! G d9„ d:e¦  «        ¦   «         ZF G d;„ d<eF¦  «        ZG e!d=¬¦  «         G d>„ d?eF¦  «        ¦   «         ZH G d@„ dAe	j8        ¦  «        ZI G dB„ dCe	j8        ¦  «        ZJ G dD„ dEe	j8        ¦  «        ZK G dF„ dGe	j8        ¦  «        ZL G dH„ dIe¦  «        ZM G dJ„ dKe	j8        ¦  «        ZN G dL„ dMe	jO        ¦  «        ZP G dN„ dOe	j8        ¦  «        ZQ e!dP¬¦  «         G dQ„ dReF¦  «        ¦   «         ZRg dS¢ZSdS )Vé    N)ÚCallable)Ú	dataclass)ÚTensoré   )Úinitialization)ÚACT2FN)ÚGradientCheckpointingLayer)ÚBaseModelOutputÚBaseModelOutputWithPooling)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)Ú#compile_compatible_method_lru_cache)ÚModelOutputÚauto_docstringÚcan_return_tupleÚlogging)ÚTransformersKwargsÚis_flash_attention_requestedÚmerge_with_config_defaults)ÚOutputRecorderÚcapture_outputsé   )Ú	AutoModelé   )Ú
Sam2ConfigÚSam2HieraDetConfigÚSam2MaskDecoderConfigÚSam2PromptEncoderConfigÚSam2VisionConfigz,Base class for the vision encoder's outputs.)Úcustom_introc                   óP   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dS )ÚSam2VisionEncoderOutputaÍ  
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, height, width, hidden_size)`):
        Sequence of hidden-states at the output of the last layer of the model.
    fpn_hidden_states (`tuple(torch.FloatTensor)`):
        Tuple of `torch.FloatTensor` (one for each feature level, from high to low resolution) of shape
        `(batch_size, hidden_size, height, width)`. Feature maps from the Feature Pyramid Network neck.
    fpn_position_encoding (`tuple(torch.FloatTensor)`):
        Tuple of `torch.FloatTensor` (one for each feature level, from high to low resolution) of shape
        `(batch_size, hidden_size, height, width)`. Positional encodings corresponding to the `fpn_hidden_states`.
    NÚfpn_hidden_statesÚfpn_position_encoding)	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r$   ÚtorchÚFloatTensorÚ__annotations__r%   © ó    úd/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/sam2/modeling_sam2.pyr#   r#   6   sP   € € € € € € ð	ð 	ð 37Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6Ø6:Ð˜5Ô,¨tÑ3Ð:Ð:Ñ:Ð:Ð:r.   r#   z'Base class for the Sam2 model's output.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         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Zeej        df         dz  ed
<   dS )ÚSam2ImageSegmentationOutputaÉ  
    iou_scores (`torch.FloatTensor` of shape `(batch_size, point_batch_size, num_masks)`):
        The Intersection over Union (IoU) scores of the predicted masks.
    pred_masks (`torch.FloatTensor` of shape `(batch_size, point_batch_size, num_masks, height, width)`):
        The predicted low-resolution masks. This is an alias for `low_res_masks`. These masks need to be post-processed
        by the processor to be brought to the original image size.
    object_score_logits (`torch.FloatTensor` of shape `(batch_size, point_batch_size, 1)`):
        Logits for the object score, indicating if an object is present.
    image_embeddings (`tuple(torch.FloatTensor)`):
        The features from the FPN, which are used by the mask decoder. This is a tuple of `torch.FloatTensor` where each
        tensor has shape `(batch_size, channels, height, width)`.
    vision_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of each stage) of shape `(batch_size, height, width, hidden_size)`.
        Hidden-states of the vision model at the output of each stage.
    vision_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length, sequence_length)`.
        Attentions weights of the vision model.
    mask_decoder_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length, sequence_length)`.
        Attentions weights of the mask decoder.
    NÚ
iou_scoresÚ
pred_masksÚobject_score_logits.Úimage_embeddingsÚvision_hidden_statesÚvision_attentionsÚmask_decoder_attentions)r&   r'   r(   r)   r2   r*   r+   r,   r3   r4   r5   Útupler6   r7   r8   r-   r.   r/   r1   r1   H   sî   € € € € € € ðð ð, ,0€J�Ô! DÑ(Ð/Ð/Ñ/Ø+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø48Ð˜Ô*¨TÑ1Ð8Ð8Ñ8Ø6:Ð�e˜EÔ-¨sÐ2Ô3Ð:Ð:Ñ:ØAEÐ˜% Ô 1°3Ð 6Ô7¸$Ñ>ÐEÐEÑEØ>BÐ�u˜UÔ.°Ð3Ô4°tÑ;ÐBÐBÑBØDHÐ˜U 5Ô#4°cÐ#9Ô:¸TÑAÐHÐHÑHÐHÐHr.   r1   c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )ÚSam2PatchEmbeddingsaê  
    Turns pixel values into patch embeddings for transformer consumption.

    Args:
        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using
            [`AutoImageProcessor`]. See [`Sam2ImageProcessor.__call__`] for details.

    Returns:
        embeddings (`torch.FloatTensor`):
            Patch embeddings depend on image_size, patch_kernel_size, patch_stride and patch_padding
    Úconfigc                 ó¾   •— t          ¦   «                              ¦   «          |j        }|j        }t	          j        |||j        |j        |j        ¬¦  «        | _	        d S )N)Úkernel_sizeÚstrideÚpadding)
ÚsuperÚ__init__Únum_channelsÚhidden_sizeÚnnÚConv2dÚpatch_kernel_sizeÚpatch_strideÚpatch_paddingÚ
projection)Úselfr<   rC   rD   Ú	__class__s       €r/   rB   zSam2PatchEmbeddings.__init__x   s]   ø€ Ý‰Œ×ÒÑÔÐØÔ*ˆØÔ(ˆåœ)ØØØÔ0ØÔ&ØÔ(ð
ñ 
ô 
ˆŒˆˆr.   c                 ó¸   — |j         \  }}}}|                      |                     | j        j        j        ¦  «        ¦  «                             dddd¦  «        }|S )Nr   r   r   r   )ÚshaperJ   ÚtoÚweightÚdtypeÚpermute)rK   Úpixel_valuesÚ_rC   ÚheightÚwidthÚ
embeddingss          r/   ÚforwardzSam2PatchEmbeddings.forward…   sV   € Ø)5Ô);Ñ&ˆˆ<˜ Ø—_’_ \§_¢_°T´_Ô5KÔ5QÑ%RÔ%RÑSÔS×[Ò[Ð\]Ð_`ÐbcÐefÑgÔgˆ
ØÐr.   )r&   r'   r(   r)   r   rB   rX   Ú__classcell__©rL   s   @r/   r;   r;   j   s^   ø€ € € € € ðð ð
Ð1ð 
ð 
ð 
ð 
ð 
ð 
ðð ð ð ð ð ð r.   r;   c                   óP  ‡ — e Zd ZdZ	 	 	 	 ddededed	edz  fˆ fd
„Ze e	d¬¦  «        	 	 	 	 dde
j        de
j        ez  de
j        deded	edz  dede
j        dz  de
j        fd„¦   «         ¦   «         Z	 dde
j        de
j        ez  de
j        de
j        dz  de
j        f
d„Zˆ xZS )ÚSam2SinePositionEmbeddingz¬
    This is a more standard version of the position embedding, very similar to the one used by the Attention is all you
    need paper, generalized to work on images.
    é@   é'  FNÚnum_position_featuresÚtemperatureÚ	normalizeÚscalec                 óÌ   •— t          ¦   «                              ¦   «          |�|du rt          d¦  «        ‚|| _        || _        || _        |€dt          j        z  n|| _        d S )NFz+normalize should be True if scale is passedr   )	rA   rB   Ú
ValueErrorr_   r`   ra   ÚmathÚpirb   )rK   r_   r`   ra   rb   rL   s        €r/   rB   z"Sam2SinePositionEmbedding.__init__‘   sj   ø€ õ 	‰Œ×ÒÑÔÐØÐ ¨eÐ!3Ð!3ÝÐJÑKÔKÐKØ%:ˆÔ"Ø&ˆÔØ"ˆŒØ$) M�Q�œ‘[�[°uˆŒ
ˆ
ˆ
r.   r   )ÚmaxsizerN   ÚdevicerQ   ÚmaskÚreturnc           
      ó   — | \  }}	}
}|€wt          j        d|
dz   ||¬¦  «        d d d …d f                              ||
|¦  «        }t          j        d|dz   ||¬¦  «        d d d d …f                              ||
|¦  «        }n?|                     |¦  «        }|                     d¦  «        }|                     d¦  «        }|r6d}||d d …dd …d d …f         |z   z  |z  }||d d …d d …dd …f         |z   z  |z  }t          j        |t           j        |¬¦  «                             |¦  «        }|dt          j        |dd¬¦  «        z  |z  z  }|d d …d d …d d …d f         |z  }|d d …d d …d d …d f         |z  }t          j        |d d …d d …d d …dd d…f                              ¦   «         |d d …d d …d d …dd d…f          	                    ¦   «         fd	¬
¦  «         
                    d¦  «        }t          j        |d d …d d …d d …dd d…f                              ¦   «         |d d …d d …d d …dd d…f          	                    ¦   «         fd	¬
¦  «         
                    d¦  «        }t          j        ||fd¬
¦  «                             dddd¦  «        }|S )Nr   ©rQ   rh   r   ç�íµ ÷Æ°>éÿÿÿÿÚfloor)Úrounding_moder   é   ©Údimr   )r*   ÚarangeÚexpandrO   ÚcumsumÚint64ÚdivÚstackÚsinÚcosÚflattenÚcatrR   )rN   rh   rQ   r_   ra   rb   r`   ri   Ú
batch_sizerT   rU   rV   Úy_embedÚx_embedÚ
embed_maskÚepsÚdim_tÚpos_xÚpos_yÚposs                       r/   Úbuild_sine_position_embeddingz7Sam2SinePositionEmbedding.build_sine_position_embedding    s  € ð (-Ñ$ˆ
�A�v˜uØˆ<õ ”l 1 f¨q¡j¸ÀfÐMÑMÔMÈdÐTUÐTUÐTUÐW[ÈmÔ\×cÒcØ˜F Eñô ˆGõ ”l 1 e¨a¡i°uÀVÐLÑLÔLÈTÐSWÐYZÐYZÐYZÈ]Ô[×bÒbØ˜F Eñô ˆGˆGð Ÿš ™œˆJØ ×'Ò'¨Ñ*Ô*ˆGØ ×'Ò'¨Ñ*Ô*ˆGØð 	CØˆCØ ¨¨¨¨B¨C¨C°°°¨Ô!3°cÑ!9Ñ:¸UÑBˆGØ ¨¨¨¨A¨A¨A¨r¨s¨s¨Ô!3°cÑ!9Ñ:¸UÑBˆGå”Ð2½%¼+ÈfÐUÑUÔU×XÒXÐY^Ñ_Ô_ˆØ ¥E¤I¨e°QÀgÐ$NÑ$NÔ$NÑ NÐQfÑ fÑgˆà˜˜˜˜1˜1˜1˜a˜a˜a ˜Ô&¨Ñ.ˆØ˜˜˜˜1˜1˜1˜a˜a˜a ˜Ô&¨Ñ.ˆÝ”˜U 1 1 1 a a a¨¨¨¨A¨D¨q¨D =Ô1×5Ò5Ñ7Ô7¸¸q¸q¸qÀ!À!À!ÀQÀQÀQÈÈÈ1È¸}Ô9M×9QÒ9QÑ9SÔ9SÐTÐZ[Ð\Ñ\Ô\×dÒdÐefÑgÔgˆÝ”˜U 1 1 1 a a a¨¨¨¨A¨D¨q¨D =Ô1×5Ò5Ñ7Ô7¸¸q¸q¸qÀ!À!À!ÀQÀQÀQÈÈÈ1È¸}Ô9M×9QÒ9QÑ9SÔ9SÐTÐZ[Ð\Ñ\Ô\×dÒdÐefÑgÔgˆÝŒi˜ ˜¨AÐ.Ñ.Ô.×6Ò6°q¸!¸QÀÑBÔBˆØˆ
r.   c           
      ób   — |                       |||| j        | j        | j        | j        |¦  «        S ©N)r‡   r_   ra   rb   r`   )rK   rN   rh   rQ   ri   s        r/   rX   z!Sam2SinePositionEmbedding.forwardÌ   s9   € ð ×1Ò1Ø�6˜5 $Ô"<¸d¼nÈdÌjÐZ^ÔZjÐlpñ
ô 
ð 	
r.   )r]   r^   FN)FNr^   Nr‰   )r&   r'   r(   r)   ÚintÚboolÚfloatrB   Ústaticmethodr   r*   ÚSizerh   ÚstrrQ   r   r‡   rX   rY   rZ   s   @r/   r\   r\   ‹   s¡  ø€ € € € € ðð ð &(Ø ØØ"ð=ð =à"ð=ð ð=ð ð	=ð
 �t‰|ð=ð =ð =ð =ð =ð =ð Ø(Ð(°Ð3Ñ3Ô3ð  Ø"Ø Ø$(ð(ð (ØŒzð(à”˜sÑ"ð(ð Œ{ð(ð  #ð	(ð
 ð(ð �t‰|ð(ð ð(ð Œl˜TÑ!ð(ð 
Œð(ð (ð (ñ 4Ô3ñ „\ð(ð^ %)ð	
ð 	
àŒzð	
ð ”˜sÑ"ð	
ð Œ{ð		
ð
 Œl˜TÑ!ð	
ð 
Œð	
ð 	
ð 	
ð 	
ð 	
ð 	
ð 	
ð 	
r.   r\   c                   ó‚   ‡ — e Zd Zdefˆ fd„Zdej        deeej        df         eej        df         f         fd„Zˆ xZ	S )ÚSam2VisionNeckr<   c           
      óx  •— t          ¦   «                              ¦   «          || _        t          |j        dz  d¬¦  «        | _        t          j        ¦   «         | _        |j	        D ]G}| j         
                    t          j        ||j        |j        |j        |j        ¬¦  «        ¦  «         ŒH|j        | _        d S )Nr   T)r_   ra   )Úin_channelsÚout_channelsr>   r?   r@   )rA   rB   r<   r\   Úfpn_hidden_sizeÚposition_encodingrE   Ú
ModuleListÚconvsÚbackbone_channel_listÚappendrF   Úfpn_kernel_sizeÚ
fpn_strideÚfpn_paddingÚfpn_top_down_levels)rK   r<   r“   rL   s      €r/   rB   zSam2VisionNeck.__init__Ù   sÆ   ø€ Ý‰Œ×ÒÑÔÐØˆŒå!:Ø"(Ô"8¸AÑ"=Èð"
ñ "
ô "
ˆÔõ ”]‘_”_ˆŒ
Ø!Ô7ð 		ð 		ˆKØŒJ×ÒÝ”	Ø +Ø!'Ô!7Ø &Ô 6Ø!Ô,Ø"Ô.ðñ ô ñô ð ð ð $*Ô#=ˆÔ Ð Ð r.   Úhidden_statesrj   .c                 óŠ  — d}d}t          | j        ¦  «        dz
  }t          |dd¦  «        D �]}||                              dddd¦  «        } | j        ||z
           |                     | j        |         j        j        ¦  «        ¦  «        }|| j        vs||k    r|}nTt          j	        |                     t          j        ¬¦  «        dd	d d
¬¦  «                             |j        ¦  «        }||z   }|                      |j        |j        |j        ¦  «                             |j        ¦  «        }	||fz  }||	fz  }�Œ||fS )Nr-   r   rn   r   r   r   )rQ   g       @ÚnearestF)Úscale_factorÚmodeÚalign_cornersÚ	antialias)Úlenr˜   ÚrangerR   rO   rP   rQ   rž   ÚFÚinterpolater*   Úfloat32r–   rN   rh   )
rK   rŸ   r$   r%   ÚnÚiÚlateral_featuresÚprev_featuresÚtop_down_featuresÚprev_position_encodings
             r/   rX   zSam2VisionNeck.forwardí   sq  € ØÐØ "Ðõ �”
‰OŒO˜aÑˆÝ�q˜"˜bÑ!Ô!ð 	?ñ 	?ˆAØ,¨QÔ/×7Ò7¸¸1¸aÀÑCÔCÐØ0˜tœz¨!¨a©%Ô0Ð1A×1DÒ1DÀTÄZÐPQÄ]ÔEYÔE_Ñ1`Ô1`ÑaÔaÐØ˜Ô0Ð0Ð0°A¸²F°FØ 0��å$%¤MØ!×$Ò$­5¬=Ð$Ñ9Ô9Ø!$Ø"Ø"&Ø#ð%ñ %ô %÷ ’"Ð%Ô+Ñ,Ô,ð "ð !1Ð3DÑ D�à%)×%;Ò%;ØÔ# ]Ô%9¸=Ô;Nñ&ô &çŠb�Ô$Ñ%Ô%ð #ð  -Ð!1Ñ1ÐØ!Ð&<Ð%>Ñ>Ð!Ñ!à Ð"7Ð7Ð7r.   )
r&   r'   r(   r    rB   r*   r   r9   rX   rY   rZ   s   @r/   r‘   r‘   Ø   s�   ø€ € € € € ð>Ð/ð >ð >ð >ð >ð >ð >ð(8 U¤\ð 8°e¸EÀ%Ä,ÐPSÐBSÔ<TÐV[Ð\aÔ\hÐjmÐ\mÔVnÐ<nÔ6oð 8ð 8ð 8ð 8ð 8ð 8ð 8ð 8r.   r‘   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingÚdropoutc                 óÀ  — t          j        ||                     dd¦  «        ¦  «        |z  }|�||z   }t          j                             |dt           j        ¬¦  «                             |j        ¦  «        }t          j         	                    ||| j
        ¬¦  «        }t          j        ||¦  «        }	|	                     dd¦  «                             ¦   «         }	|	|fS )Nr   r   rn   )rs   rQ   )ÚpÚtrainingr   )r*   ÚmatmulÚ	transposerE   Ú
functionalÚsoftmaxrª   rO   rQ   r¸   r»   Ú
contiguous)
r²   r³   r´   rµ   r¶   r·   r¸   ÚkwargsÚattn_weightsÚattn_outputs
             r/   Úeager_attention_forwardrÄ     sÃ   € õ ”<  s§}¢}°Q¸Ñ':Ô':Ñ;Ô;¸gÑE€LØÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐW\ÔWbÑcÔc€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€LÝ”,˜|¨UÑ3Ô3€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r.   ÚxÚquery_striderj   c                 ó´   — |€| S |                       dddd¦  «        } t          j                             | ||d¬¦  «        } |                       dddd¦  «        } | S )Nr   r   r   r   F)r>   r?   Ú	ceil_mode)rR   rE   r¾   Ú
max_pool2d)rÅ   rÆ   s     r/   Údo_poolrÊ   "  s_   € ØÐØˆà	�	Š	�!�Q˜˜1ÑÔ€AÝ
Œ× Ò  °À\Ð]bÐ ÑcÔc€Aà	�	Š	�!�Q˜˜1ÑÔ€AØ€Hr.   c                   ór   ‡ — e Zd Z	 ddededededeeef         dz  f
ˆ fd„Zdej        d	ej        fd
„Z	ˆ xZ
S )ÚSam2MultiScaleAttentionNr<   rs   Údim_outÚnum_attention_headsrÆ   c                 ó(  •— t          ¦   «                              ¦   «          || _        || _        || _        || _        || _        ||z  }|dz  | _        t          j	        ||dz  ¦  «        | _
        t          j	        ||¦  «        | _        d| _        d S )Nç      à¿r   F)rA   rB   r<   rs   rÍ   rÆ   rÎ   rb   rE   ÚLinearÚqkvÚprojÚ	is_causal)rK   r<   rs   rÍ   rÎ   rÆ   Úhead_dimrL   s          €r/   rB   z Sam2MultiScaleAttention.__init__.  s‹   ø€ õ 	‰Œ×ÒÑÔÐàˆŒàˆŒØˆŒØ(ˆÔà#6ˆÔ ØÐ1Ñ1ˆØ˜t‘^ˆŒ
Ý”9˜S '¨A¡+Ñ.Ô.ˆŒÝ”I˜g wÑ/Ô/ˆŒ	àˆŒˆˆr.   rŸ   rj   c                 ó´  — |j         \  }}}}|                      |¦  «                             |||z  d| j        d¦  «        }t	          j        |d¦  «        \  }}	}
|| j        z  |	                     dd¦  «        z  }t          j        j	         
                    |t          j        d¬¦  «                             |j        ¦  «        }| j        r]t          |                     |||d¦  «        | j        ¦  «        }|j         dd…         \  }}|                     |||z  | j        d¦  «        }|                     dd¦  «        }|	                     dd¦  «        }	|
                     dd¦  «        }
t!          j        | j        j        t(          ¦  «        } || ||	|
fd | j        | j        dœ|¤Ž\  }}|                     |||d¦  «        }|                      |¦  «        }|S )Nr   rn   r   éþÿÿÿ)rQ   rs   r   )r¶   rÔ   r·   )rN   rÒ   ÚreshaperÎ   r*   Úunbindrb   r½   rE   r¾   r¿   rª   rO   rQ   rÆ   rÊ   r   Úget_interfacer<   Ú_attn_implementationrÄ   rÔ   rÓ   )rK   rŸ   rÁ   r~   rU   rV   rT   rÒ   r³   r´   rµ   rÂ   Úattention_interfacerÃ   s                 r/   rX   zSam2MultiScaleAttention.forwardF  sä  € Ø'4Ô':Ñ$ˆ
�F˜E 1à�hŠh�}Ñ%Ô%×-Ò-¨j¸&À5¹.È!ÈTÔMeÐgiÑjÔjˆå!œL¨¨aÑ0Ô0Ñˆˆs�Eà ¤
Ñ*¨c¯mªm¸BÀÑ.CÔ.CÑCˆÝ”xÔ*×2Ò2°<ÅuÄ}ÐZ\Ð2Ñ]Ô]×`Ò`ÐafÔalÑmÔmˆð Ôð 	\Ý˜EŸMšM¨*°f¸eÀRÑHÔHÈ$ÔJ[Ñ\Ô\ˆEØ!œK¨¨!¨Ô,‰MˆF�EØ—M’M *¨f°u©n¸dÔ>VÐXZÑ[Ô[ˆEð —’  1Ñ%Ô%ˆØ�mŠm˜A˜qÑ!Ô!ˆØ—’  1Ñ%Ô%ˆå(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð -Ð,ØØØØð		
ð
  Ø”nØ”Jð	
ð 	
ð ð	
ð 	
‰ˆ�Qð "×)Ò)¨*°f¸eÀRÑHÔHˆà—i’i Ñ,Ô,ˆàÐr.   r‰   )r&   r'   r(   r   rŠ   r9   rB   r*   r   rX   rY   rZ   s   @r/   rÌ   rÌ   -  sª   ø€ € € € € ð 04ðð à"ðð ðð ð	ð
 !ðð ˜C ˜H”o¨Ñ,ðð ð ð ð ð ð0& U¤\ð &ÀÄð &ð &ð &ð &ð &ð &ð &ð &r.   rÌ   c                   óD   ‡ — e Zd Z	 	 ddedededededefˆ fd	„Zd
„ Zˆ xZS )ÚSam2FeedForwardÚreluFÚ	input_dimÚ
hidden_dimÚ
output_dimÚ
num_layersÚ
activationÚsigmoid_outputc                 ó\  •‡— t          ¦   «                              ¦   «          || _        t          |         | _        t          j        |‰¦  «        | _        t          j        ‰|¦  «        | _        t          j	        ˆfd„t          |dz
  ¦  «        D ¦   «         ¦  «        | _        || _        d S )Nc                 ó:   •— g | ]}t          j        ‰‰¦  «        ‘ŒS r-   )rE   rÑ   )Ú.0rT   rá   s     €r/   ú
<listcomp>z,Sam2FeedForward.__init__.<locals>.<listcomp>~  s%   ø€ Ð$fÐ$fÐ$fÈ1¥R¤Y¨z¸:Ñ%FÔ%FÐ$fÐ$fÐ$fr.   r   )rA   rB   rã   r   rä   rE   rÑ   Úproj_inÚproj_outr—   r§   Úlayersrå   )rK   rà   rá   râ   rã   rä   rå   rL   s     `    €r/   rB   zSam2FeedForward.__init__p  s™   øø€ õ 	‰Œ×ÒÑÔÐØ$ˆŒÝ  Ô,ˆŒÝ”y ¨JÑ7Ô7ˆŒÝœ	 *¨jÑ9Ô9ˆŒÝ”mÐ$fÐ$fÐ$fÐ$fÕPUÐV`ÐcdÑVdÑPeÔPeÐ$fÑ$fÔ$fÑgÔgˆŒØ,ˆÔÐÐr.   c                 ó
  — |                       |¦  «        }|                      |¦  «        }| j        D ] }|                       ||¦  «        ¦  «        }Œ!|                      |¦  «        }| j        rt          j        |¦  «        }|S r‰   )rê   rä   rì   rë   rå   r¨   Úsigmoid)rK   rŸ   Úlayers      r/   rX   zSam2FeedForward.forward�  s…   € ØŸš ]Ñ3Ô3ˆØŸš¨Ñ6Ô6ˆØ”[ð 	Bð 	BˆEØ ŸOšO¨E¨E°-Ñ,@Ô,@ÑAÔAˆMˆMàŸš mÑ4Ô4ˆØÔð 	5ÝœI mÑ4Ô4ˆMØÐr.   )rß   F)	r&   r'   r(   rŠ   r�   r‹   rB   rX   rY   rZ   s   @r/   rÞ   rÞ   o  s“   ø€ € € € € ð !Ø$ð-ð -àð-ð ð-ð ð	-ð
 ð-ð ð-ð ð-ð -ð -ð -ð -ð -ð"	ð 	ð 	ð 	ð 	ð 	ð 	r.   rÞ   c           	      ój  — | j         \  }}}}| |z  }| |z  }t          j                             | ddd|d|f¦  «        } ||z   ||z   }	}||z  }
|	|z  }|                      ||
||||¦  «        } |                      dddddd¦  «                             ¦   «                              d|||¦  «        }|||	ffS )a  
    Partition into non-overlapping windows with padding if needed.

    Args:
        hidden_state (`torch.Tensor`):
            Input tokens with [batch_size, height, width, num_channels].
        window_size (`int`):
            Window size.

    Returns:
        `tuple(torch.FloatTensor)` comprising various elements:
        - windows: windows after partition with [batch_size * num_windows, window_size, window_size, num_channels].
        - (padded_height, padded_width): padded height and width before partition
    r   r   r   r   rq   é   rn   )rN   rE   r¾   ÚpadÚviewrR   rÀ   )Úhidden_stateÚwindow_sizer~   rU   rV   rC   Ú
pad_heightÚ	pad_widthÚpadded_heightÚpadded_widthÚn_hÚn_wÚwindowss                r/   Úwindow_partitionrý   �  sç   € ð /;Ô.@Ñ+€J�˜˜|ð �'˜[Ñ(€JØ�˜;Ñ&€IÝ”=×$Ò$ \°A°q¸!¸YÈÈ:Ð3VÑWÔW€LØ"(¨:Ñ"5°u¸yÑ7H�<€Mà
˜;Ñ
&€CØ
˜+Ñ
%€CØ×$Ò$ Z°°kÀ3ÈÐUaÑbÔb€LØ×"Ò" 1 a¨¨A¨q°!Ñ4Ô4×?Ò?ÑAÔA×FÒFÀrÈ;ÐXcÐeqÑrÔr€GØ�] LÐ1Ð1Ð1r.   c                 óX  — |\  }}|\  }}||z  }||z  }	| j         d         ||	z  z  }
|                      |
||	||d¦  «        }|                     dddddd¦  «                             ¦   «         }|                     |
||d¦  «        }|dd…d|…d|…dd…f                              ¦   «         S )	aB  
    Window unpartition into original sequences and removing padding.

    Args:
        windows (`torch.Tensor`):
            Input tokens with [batch_size * num_windows, window_size, window_size, num_channels].
        window_size (`int`):
            Window size.
        pad_height_width (`tuple[int]`):
            Padded height and width (padded_height, padded_width).
        height_width (`tuple[int]`):
            Original height and width before padding.

    Returns:
        hidden_state: unpartitioned sequences with [batch_size, height, width, num_channels].
    r   rn   r   r   r   rq   rñ   N)rN   ró   rR   rÀ   )rü   rõ   Úpad_height_widthÚheight_widthrø   rù   rU   rV   rú   rû   r~   rô   s               r/   Úwindow_unpartitionr  ¬  sÎ   € ð" #3Ñ€M�<Ø �M€FˆEØ
˜;Ñ
&€CØ
˜+Ñ
%€CØ”˜qÔ! c¨C¡iÑ0€JØ—<’< 
¨C°°kÀ;ÐPRÑSÔS€LØ×'Ò'¨¨1¨a°°A°qÑ9Ô9×DÒDÑFÔF€LØ×$Ò$ Z°ÀÈbÑQÔQ€Là˜˜˜˜7˜F˜7 F U F¨A¨A¨AÐ-Ô.×9Ò9Ñ;Ô;Ð;r.   c                   ód   ‡ — e Zd Zdedededefˆ fd„Zdej        dee	         dej
        fd	„Zˆ xZS )
ÚSam2MultiScaleBlockr<   Ú	stage_idxÚ	block_idxÚtotal_block_idxc                 óŽ  •— t          ¦   «                              ¦   «          |dk    r|dk    r|j        |dz
           n|j        |         | _        |j        |         | _        t          j        | j        |j        ¬¦  «        | _        |dk    r|dk    r|j	        |dz
           n|j	        |         | _
        ||j        v rdn| j
        | _
        d|cxk     r|j        k    rn n|dk    r|j        nd | _        t          || j        | j        |j        |         | j        ¬¦  «        | _        t          j        | j        |j        ¬¦  «        | _        t%          | j        t'          | j        |j        z  ¦  «        | j        d|j        ¬¦  «        | _        | j        | j        k    r&t          j        | j        | j        ¦  «        | _        d S d S )Nr   r   )r‚   )rÎ   rÆ   r   )rã   rä   )rA   rB   Úembed_dim_per_stagers   rÍ   rE   Ú	LayerNormÚlayer_norm_epsÚlayer_norm1Úwindow_size_per_stagerõ   Úglobal_attention_blocksÚnum_query_pool_stagesrÆ   rÌ   Únum_attention_heads_per_stageÚattnÚlayer_norm2rÞ   rŠ   Ú	mlp_ratioÚ
hidden_actÚmlprÑ   rÓ   )rK   r<   r  r  r  rL   s        €r/   rB   zSam2MultiScaleBlock.__init__Ê  sæ  ø€ õ 	‰Œ×ÒÑÔÐð
 ˜1Š}ˆ} ¨a¢ ð Ô& y°1¡}Ô5Ð5àÔ+¨IÔ6ð 	Œð
 Ô1°)Ô<ˆŒÝœ<¨¬°fÔ6KÐLÑLÔLˆÔð ˜1Š}ˆ} ¨a¢ ð Ô(¨°Q©Ô7Ð7àÔ-¨iÔ8ð 	Ôð
 !0°6Ô3QÐ QÐ Q˜1˜1ÐW[ÔWgˆÔð $% yÐ#PÐ#PÒ#PÐ#P°FÔ4PÒ#PÐ#PÐ#PÐ#PÐ#PÐU^ÐbcÒUcÐUcˆFÔÐÐimð 	Ôõ ,ØØŒHØŒLØ &Ô DÀYÔ OØÔ*ð
ñ 
ô 
ˆŒ	õ œ<¨¬¸&Ô:OÐPÑPÔPˆÔÝ"ØŒLÝ�”˜vÔ/Ñ/Ñ0Ô0ØŒLØØÔ(ð
ñ 
ô 
ˆŒð Œ8�t”|Ò#Ð#Ýœ	 $¤(¨D¬LÑ9Ô9ˆDŒIˆIˆIð $Ð#r.   rŸ   rÁ   rj   c                 ón  — |}|                       |¦  «        }| j        | j        k    r(t          |                      |¦  «        | j        ¦  «        }| j        }| j        dk    r-|j        d         |j        d         }}t          ||¦  «        \  }} | j	        dd|i|¤Ž}|}| j        r=| j        | j        d         z  }|j        dd…         \  }}| |z  }	| |z  }
||	z   ||
z   f}| j        dk    rt          |||||f¦  «        }||z   }|                      |¦  «        }||                      |¦  «        z   }|S )Nr   r   r   rŸ   r   r-   )r  rs   rÍ   rÊ   rÓ   rÆ   rõ   rN   rý   r  r  r  r  )rK   rŸ   rÁ   Úresidualrõ   ÚHÚWÚpad_hwrÃ   Úpad_hÚpad_wÚlayernorm_outputs               r/   rX   zSam2MultiScaleBlock.forwardù  s„  € ð
 !ˆà×(Ò(¨Ñ7Ô7ˆð Œ8�t”|Ò#Ð#Ý˜tŸyšy¨Ñ7Ô7¸Ô9JÑKÔKˆHð Ô&ˆØÔ˜aÒÐØ Ô& qÔ)¨=Ô+>¸qÔ+AˆqˆAÝ$4°]ÀKÑ$PÔ$PÑ!ˆM˜6ð  �d”ið 
ð 
Ø'ð
àð
ð 
ˆð $ˆØÔð 	,àÔ*¨dÔ.?ÀÔ.BÑBˆKØ”> ! A #Ô&‰DˆAˆqØ�R˜;Ñ&ˆEØ�R˜;Ñ&ˆEØ˜%‘i  U¡Ð+ˆFð Ô˜aÒÐÝ.¨}¸kÈ6ÐTUÐWXÐSYÑZÔZˆMà  =Ñ0ˆØ×+Ò+¨MÑ:Ô:ÐØ%¨¯ªÐ1AÑ(BÔ(BÑBˆàÐr.   )r&   r'   r(   r   rŠ   rB   r*   r   r   r   r+   rX   rY   rZ   s   @r/   r  r  É  sŸ   ø€ € € € € ð-:à"ð-:ð ð-:ð ð	-:ð
 ð-:ð -:ð -:ð -:ð -:ð -:ð^)à”|ð)ð Ð+Ô,ð)ð 
Ô	ð	)ð )ð )ð )ð )ð )ð )ð )r.   r  zW
    Hiera model's outputs that also contains a pooling of the last hidden states.
    c                   ó¼   — e Zd ZU dZ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Ze	ej        df         dz  ed<   dS )ÚSam2HieraDetModelOutputat  
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, height, width, hidden_size)`):
        hidden-states at the output of the last layer of the model.
    intermediate_hidden_states (`tuple[torch.FloatTensor]` of shape `(batch_size, height, width, hidden_size)`):
        Sequence of hidden-states at the output of the intermediate layers of the model.
    NÚlast_hidden_state.Úintermediate_hidden_statesrŸ   Ú
attentions)r&   r'   r(   r)   r  r*   r+   r,   r   r9   rŸ   r!  r-   r.   r/   r  r  %  sž   € € € € € € ðð ð 37Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6ØGKÐ  eÔ&7¸Ð&<Ô =ÀÑ DÐKÐKÑKØ:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>Ø7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r.   r  c                   ól   ‡ — e Zd ZeZdZdZdZdZdZ	dZ
g d¢Z ej        ¦   «         ˆ fd„¦   «         Zˆ xZS )ÚSam2PreTrainedModelÚsam2rS   )ÚimageT)z
^memory_.*z^mask_downsample.*z^object_pointer_proj.*z0^temporal_positional_encoding_projection_layer.*Úno_memory_positional_encodingÚno_object_pointerÚ%occlusion_spatial_embedding_parameterc                 óÜ  •— t          ¦   «                              |¦  «         t          |t          ¦  «        rD|j        �t          j        |j        ¦  «         |j        �t          j        |j        ¦  «         d S d S t          |t          ¦  «        r"t          j	        |j
        |j        ¬¦  «         d S t          |t          ¦  «        r"|j        �t          j        |j        ¦  «         d S d S d S )N)Ústd)rA   Ú_init_weightsÚ
isinstanceÚSam2HieraDetModelÚ	pos_embedÚinitÚzeros_Úpos_embed_windowÚSam2PositionalEmbeddingÚnormal_Úpositional_embeddingrb   Ú	Sam2ModelÚno_memory_embedding)rK   r²   rL   s     €r/   r+  z!Sam2PreTrainedModel._init_weightsL  së   ø€ å‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ/Ñ0Ô0ð 		8ØÔÐ+Ý”˜FÔ,Ñ-Ô-Ð-ØÔ&Ð2Ý”˜FÔ3Ñ4Ô4Ð4Ð4Ð4ð 3Ð2å˜Õ 7Ñ8Ô8ð 	8ÝŒL˜Ô4¸&¼,ÐGÑGÔGÐGÐGÐGÝ˜¥	Ñ*Ô*ð 	8ØÔ)Ð5Ý”˜FÔ6Ñ7Ô7Ð7Ð7Ð7ð	8ð 	8Ø5Ð5r.   )r&   r'   r(   r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚinput_modalitiesÚ_supports_sdpaÚ_supports_flash_attnÚ_supports_attention_backendÚ"_keys_to_ignore_on_load_unexpectedr*   Úno_gradr+  rY   rZ   s   @r/   r#  r#  9  s…   ø€ € € € € à€LØÐØ$€OØ!ÐØ€NØÐØ"&Ðð*ð *ð *Ð&ð €U„]�_„_ð8ð 8ð 8ð 8ñ „_ð8ð 8ð 8ð 8ð 8r.   r#  c            
       óÀ   ‡ — e Zd ZeZdZeedœZdefˆ fd„Z	d„ Z
deeef         dej        fd„Zee	 ddej        d	z  d
ee         deez  fd„¦   «         ¦   «         Zˆ xZS )r-  rS   ©rŸ   r!  r<   c           	      óê  •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          j        t          j        d|j        g|j	        ¢R Ž ¦  «        | _
        t	          j        t          j        d|j        |j        d         |j        d         ¦  «        ¦  «        | _        t          j        |j        ¦  «        dz
                       ¦   «         | _        t	          j        ¦   «         | _        d}t)          |j        ¦  «        D ]I\  }}t+          |¦  «        D ]4}t-          ||||¬¦  «        }| j                             |¦  «         |dz  }Œ5ŒJ|                      ¦   «          d S )Nr   r   )r<   r  r  r  )rA   rB   r;   Úpatch_embedrE   Ú	Parameterr*   ÚzerosrD   Ú+window_positional_embedding_background_sizer.  r  r1  Únprv   Úblocks_per_stageÚtolistÚ
stage_endsr—   ÚblocksÚ	enumerater§   r  rš   Ú	post_init)rK   r<   r  r  rH  r  ÚblockrL   s          €r/   rB   zSam2HieraDetModel.__init__c  si  ø€ Ý‰Œ×Ò˜Ñ Ô Ð å.¨vÑ6Ô6ˆÔåœÝŒK˜˜6Ô-Ðc°Ô0bÐcÐcÐcñ
ô 
ˆŒõ !#¤ÝŒK˜˜6Ô-¨vÔ/KÈAÔ/NÐPVÔPlÐmnÔPoÑpÔpñ!
ô !
ˆÔõ œ9 VÔ%<Ñ=Ô=ÀÑA×IÒIÑKÔKˆŒÝ”m‘o”oˆŒØˆÝ+4°VÔ5LÑ+MÔ+Mð 	%ð 	%Ñ'ˆIÐ'Ý"Ð#3Ñ4Ô4ð %ð %�	Ý+Ø!¨YÀ)Ð]lðñ ô �ð ”×"Ò" 5Ñ)Ô)Ð)Ø 1Ñ$��ð%ð 	�ŠÑÔÐÐÐr.   c                 ó   — | j         S r‰   )rC  ©rK   s    r/   Úget_input_embeddingsz&Sam2HieraDetModel.get_input_embeddings{  s   € ØÔÐr.   Úhwrj   c                 óþ   — |\  }}| j         }t          j        | j        ||fd¬¦  «        }||                     d„ t          |j        |j        ¦  «        D ¦   «         ¦  «        z   }|                     dddd¦  «        }|S )NÚbicubic)Úsizer£   c                 ó   — g | ]
\  }}||z  ‘ŒS r-   r-   )rè   rÅ   Úys      r/   ré   z4Sam2HieraDetModel._get_pos_embed.<locals>.<listcomp>‚  s    € Ð2oÐ2oÐ2o¹d¸aÀ°1¸±6Ð2oÐ2oÐ2or.   r   r   r   r   )r1  r¨   r©   r.  ÚtileÚziprN   rR   )rK   rR  ÚhÚwÚwindow_embedr.  s         r/   Ú_get_pos_embedz Sam2HieraDetModel._get_pos_embed~  s…   € Ø‰ˆˆ1ØÔ,ˆÝ”M $¤.¸¸1°vÀIÐNÑNÔNˆ	Ø × 1Ò 1Ð2oÐ2oÅcÈ)Ì/Ð[gÔ[mÑFnÔFnÐ2oÑ2oÔ2oÑ pÔ pÑpˆ	Ø×%Ò% a¨¨A¨qÑ1Ô1ˆ	ØÐr.   NrÁ   c                 ó"  — |€t          d¦  «        ‚|                      |¦  «        }||                      |j        dd…         ¦  «        z   }d}t	          | j        ¦  «        D ]\  }} ||fi |¤Ž}|| j        v r||fz   }Œt          ||¬¦  «        S )Nú You have to specify pixel_valuesr   r   r-   )r  r   )rd   rC  r]  rN   rL  rK  rJ  r  )rK   rS   rÁ   rŸ   r   r¬   Úblock_modules          r/   rX   zSam2HieraDetModel.forward†  sÈ   € ð ÐÝÐ?Ñ@Ô@Ð@à×(Ò(¨Ñ6Ô6ˆØ%¨×(;Ò(;¸MÔ<OÐPQÐRSÐPSÔ<TÑ(UÔ(UÑUˆà%'Ð"Ý(¨¬Ñ5Ô5ð 	[ð 	[‰OˆAˆ|Ø(˜L¨ÐAÐA¸&ÐAÐAˆMà�D”OÐ#Ð#Ø-GÈ=ÐJZÑ-ZÐ*øå&Ø+Ø'Að
ñ 
ô 
ð 	
r.   r‰   )r&   r'   r(   r   r7  r9  r  rÌ   Ú_can_record_outputsrB   rQ  r9   rŠ   r*   r   r]  r   r   r+   r   r   r  rX   rY   rZ   s   @r/   r-  r-  [  s  ø€ € € € € Ø%€LØ$€Oà,Ø-ðð Ðð
Ð1ð ð ð ð ð ð ð0 ð  ð  ð  s¨C x¤ð °U´\ð ð ð ð ð  Øð 26ð
ð 
àÔ'¨$Ñ.ð
ð Ð+Ô,ð
ð 
Ð(Ñ	(ð	
ð 
ð 
ñ „_ñ  Ôð
ð 
ð 
ð 
ð 
r.   r-  zJ
    The vision model from Sam without any head or projection on top.
    c            	       ó†   ‡ — e Zd ZeZdZeedœZdefˆ fd„Z	d„ Z
e	 d
dej        dz  dee         deez  fd	„¦   «         Zˆ xZS )ÚSam2VisionModelrS   rA  r<   c                 óü   •— t          ¦   «                              |¦  «         || _        t          j        |j        ¦  «        | _        t          |¦  «        | _        |j	        | _	        |  
                    ¦   «          d S r‰   )rA   rB   r<   r   Úfrom_configÚbackbone_configÚbackboner‘   ÚneckÚnum_feature_levelsrM  ©rK   r<   rL   s     €r/   rB   zSam2VisionModel.__init__­  sg   ø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒå!Ô-¨fÔ.DÑEÔEˆŒå" 6Ñ*Ô*ˆŒ	Ø"(Ô";ˆÔà�ŠÑÔÐÐÐr.   c                 ó4   — | j                              ¦   «         S r‰   )rg  rQ  rP  s    r/   rQ  z$Sam2VisionModel.get_input_embeddings¸  s   € ØŒ}×1Ò1Ñ3Ô3Ð3r.   NrÁ   rj   c                 ó,  — |€t          d¦  «        ‚ | j        |fi |¤Ž}|j        }|j        }|                      |¦  «        \  }}|| j         d …         d d d…         }|| j         d …         d d d…         }t          ||||j        |j        ¬¦  «        S )Nr_  rn   )r  r$   r%   rŸ   r!  )	rd   rg  r  r   rh  ri  r#   rŸ   r!  )rK   rS   rÁ   Úbackbone_outputrŸ   r   r$   r%   s           r/   rX   zSam2VisionModel.forward»  sÒ   € ð ÐÝÐ?Ñ@Ô@Ð@ð (˜$œ-¨Ð?Ð?¸Ð?Ð?ˆØ'Ô9ˆØ%4Ô%OÐ"à37·9²9Ð=WÑ3XÔ3XÑ0ÐÐ0à-¨tÔ/FÐ.FÐ.HÐ.HÔIÈ$È$ÈBÈ$ÔOÐØ 5°tÔ7NÐ6NÐ6PÐ6PÔ QÐRVÐRVÐTVÐRVÔ WÐå&Ø+Ø/Ø"7Ø)Ô7Ø&Ô1ð
ñ 
ô 
ð 	
r.   r‰   )r&   r'   r(   r    r7  r9  r  rÌ   ra  rB   rQ  r   r*   r+   r   r   r9   r#   rX   rY   rZ   s   @r/   rc  rc     sÎ   ø€ € € € € ð $€LØ$€Oà,Ø-ðð Ðð
	Ð/ð 	ð 	ð 	ð 	ð 	ð 	ð4ð 4ð 4ð ð 26ð
ð 
àÔ'¨$Ñ.ð
ð Ð+Ô,ð
ð 
Ð(Ñ	(ð	
ð 
ð 
ñ Ôð
ð 
ð 
ð 
ð 
r.   rc  c                   ó,   ‡ — e Zd Zdefˆ fd„Zdd„Zˆ xZS )r2  r<   c                 óØ   •— t          ¦   «                              ¦   «          |j        | _        | j        t          j        d|j        dz  f¦  «        z  }|                      d|¦  «         d S )Nr   r4  )rA   rB   rb   r*   ÚrandnrD   Úregister_buffer)rK   r<   r4  rL   s      €r/   rB   z Sam2PositionalEmbedding.__init__Ø  se   ø€ Ý‰Œ×ÒÑÔÐØ”\ˆŒ
Ø#œz­E¬K¸¸FÔ<NÐRSÑ<SÐ8TÑ,UÔ,UÑUÐØ×ÒÐ3Ð5IÑJÔJÐJÐJÐJr.   Nc                 ó
  — |                      ¦   «         }|�P|dd…dd…dd…df         |d         z  |dd…dd…dd…df<   |dd…dd…dd…df         |d         z  |dd…dd…dd…df<   |                     t          j        ¦  «         d|z  dz
  }|                     | j        j        ¦  «        }|| j        z  }dt          j        z  |z  }t          j        t          j	        |¦  «        t          j
        |¦  «        gd¬¦  «        S )z8Positionally encode points that are normalized to [0,1].Nr   r   r   rn   rr   )ÚclonerO   r*   rª   r4  rQ   rG  rf   r}   rz   r{   )rK   Úinput_coordsÚinput_shapeÚcoordinatess       r/   rX   zSam2PositionalEmbedding.forwardÞ  s  € à"×(Ò(Ñ*Ô*ˆàÐ"Ø&1°!°!°!°Q°Q°Q¸¸¸¸1°*Ô&=ÀÈAÄÑ&NˆK˜˜˜˜1˜1˜1˜a˜a˜a ˜
Ñ#Ø&1°!°!°!°Q°Q°Q¸¸¸¸1°*Ô&=ÀÈAÄÑ&NˆK˜˜˜˜1˜1˜1˜a˜a˜a ˜
Ñ#Ø�Š•u”}Ñ%Ô%Ð%ð ˜+‘o¨Ñ)ˆØ!—n’n TÔ%>Ô%DÑEÔEˆØ! DÔ$=Ñ=ˆØ�"œ%‘i +Ñ-ˆåŒy�%œ) KÑ0Ô0µ%´)¸KÑ2HÔ2HÐIÈrÐRÑRÔRÐRr.   r‰   ©r&   r'   r(   r   rB   rX   rY   rZ   s   @r/   r2  r2  ×  sh   ø€ € € € € ðKÐ6ð Kð Kð Kð Kð Kð KðSð Sð Sð Sð Sð Sð Sð Sr.   r2  c                   ó*   ‡ — e Zd Zdefˆ fd„Zd„ Zˆ xZS )ÚSam2MaskEmbeddingr<   c                 óü  •— t          ¦   «                              ¦   «          |j        dz  | _        t          |j                 | _        t          j        d| j        dd¬¦  «        | _        t          j        | j        |j        dd¬¦  «        | _	        t          j        |j        |j
        d¬¦  «        | _        t          | j        |j        d¬¦  «        | _        t          | j        dz  |j        d¬¦  «        | _        d S )Nrq   r   r   ©r>   r?   )r>   Úchannels_first©r‚   Údata_format)rA   rB   Úmask_input_channelsr   r  rä   rE   rF   Úconv1Úconv2rD   Úconv3ÚSam2LayerNormr
  r  r  rj  s     €r/   rB   zSam2MaskEmbedding.__init__ñ  sî   ø€ Ý‰Œ×ÒÑÔÐØ#)Ô#=ÀÑ#BˆÔ Ý  Ô!2Ô3ˆŒÝ”Y˜q $Ô":ÈÐRSÐTÑTÔTˆŒ
Ý”Y˜tÔ7¸Ô9SÐabÐklÐmÑmÔmˆŒ
Ý”Y˜vÔ9¸6Ô;MÐ[\Ð]Ñ]Ô]ˆŒ
Ý(ØÔ$¨&Ô*?ÐM]ð
ñ 
ô 
ˆÔõ )ØÔ$ qÑ(¨fÔ.CÐQað
ñ 
ô 
ˆÔÐÐr.   c                 ó,  — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|S r‰   )r€  r  rä   r�  r  r‚  )rK   ÚmasksrŸ   Údense_embeddingss       r/   rX   zSam2MaskEmbedding.forwardÿ  s„   € ØŸ
š
 5Ñ)Ô)ˆØ×(Ò(¨Ñ7Ô7ˆØŸš¨Ñ6Ô6ˆàŸ
š
 =Ñ1Ô1ˆØ×(Ò(¨Ñ7Ô7ˆØŸš¨Ñ6Ô6ˆØŸ:š: mÑ4Ô4ÐØÐr.   rw  rZ   s   @r/   ry  ry  ð  sT   ø€ € € € € ð
Ð6ð 
ð 
ð 
ð 
ð 
ð 
ð	 ð 	 ð 	 ð 	 ð 	 ð 	 ð 	 r.   ry  c                   ó  ‡ — e Zd Zdefˆ fd„Zdej        dej        dedej        fd„Zdej        dej        fd	„Z	d
e
ej        ej        f         dz  dej        dz  dej        dz  dej        dz  de
ej        ej        f         f
d„Zˆ xZS )ÚSam2PromptEncoderr<   c                 ó$  •— t          ¦   «                              ¦   «          t          |¦  «        | _        t	          |¦  «        | _        t          j        d|j        ¦  «        | _	        |j
        |j        z  |j
        |j        z  f| _        d|j
        z  |j        z  d|j
        z  |j        z  f| _        |j
        | _        t          j        |j        |j        ¦  «        | _        |j        | _        t          j        d|j        ¦  «        | _        d S )Nr   rq   )rA   rB   r2  Úshared_embeddingry  Ú
mask_embedrE   Ú	EmbeddingrD   Úno_mask_embedÚ
image_sizeÚ
patch_sizeÚimage_embedding_sizeÚmask_input_sizeÚinput_image_sizeÚnum_point_embeddingsÚpoint_embedÚnot_a_point_embedrj  s     €r/   rB   zSam2PromptEncoder.__init__  sï   ø€ Ý‰Œ×ÒÑÔÐÝ 7¸Ñ ?Ô ?ˆÔÝ+¨FÑ3Ô3ˆŒÝœ\¨!¨VÔ-?Ñ@Ô@ˆÔà%+Ô%6¸&Ô:KÑ%KÈVÔM^ÐbhÔbsÑMsÐ$tˆÔ!Ø ! FÔ$5Ñ 5¸Ô9JÑ JÈAÐPVÔPaÑLaÐekÔevÑLvÐwˆÔØ &Ô 1ˆÔåœ<¨Ô(CÀVÔEWÑXÔXˆÔØ!Ô-ˆÔÝ!#¤¨a°Ô1CÑ!DÔ!DˆÔÐÐr.   ÚpointsÚlabelsrò   rj   c                 ó@  — |dz   }|rPt           j        j                             |ddd¬¦  «        }t           j        j                             |ddd¬¦  «        }| j        | j        f}|                      ||¦  «        }t          j        |d         dk    | j        j        |¦  «        }t          j        |d         d	k    |t          j	        |¦  «        ¦  «        }||  
                    |                     d¬
¦  «        ¦  «        |dk                         d¦  «        z  z   }|S )zEmbeds point prompts.ç      à?©r   r   r   r   Úconstantr   ©r£   rµ   )r   r   rn   ).Niöÿÿÿ)Úmin)r*   rE   r¾   rò   r’  rŠ  Úwherer•  rP   Ú
zeros_liker”  ÚclampÚ	unsqueeze)rK   r–  r—  rò   ru  Úpoint_embeddings         r/   Ú_embed_pointszSam2PromptEncoder._embed_points  s  € à˜#‘ˆØð 	XÝ”XÔ(×,Ò,¨V°\È
ÐZ[Ð,Ñ\Ô\ˆFÝ”XÔ(×,Ò,¨V°VÀ*ÐTVÐ,ÑWÔWˆFØÔ,¨dÔ.CÐDˆØ×/Ò/°¸ÑDÔDˆõ  œ+ f¨YÔ&7¸2Ò&=¸tÔ?UÔ?\Ð^mÑnÔnˆõ  œ+Ø�9Ô Ò$ØÝÔ˜_Ñ-Ô-ñ
ô 
ˆð *¨D×,<Ò,<¸V¿\º\Èa¸\Ñ=PÔ=PÑ,QÔ,QÐU[Ð_`ÒU`×TkÒTkÐlnÑToÔToÑ,oÑoˆàÐr.   Úboxesc                 ó   — |dz   } |j         g |j        dd…         ¢d‘d‘R Ž }t          j        j                             |ddd¬¦  «        }|                      || j        | j        f¦  «        }|dd…dd…ddd…fxx         | j        j	        d         z  cc<   |dd…dd…ddd…fxx         | j        j	        d	         z  cc<   | j
        j	                             |dd…dd…ddd…f         ¦  «        |dd…dd…ddd…f<   |S )
zEmbeds box prompts.r™  Nr   rš  r›  r   rœ  r   r   )ró   rN   r*   rE   r¾   rò   rŠ  r’  r”  rP   r•  Ú	expand_as)rK   r¤  ÚcoordsÚcorner_embeddings       r/   Ú_embed_boxeszSam2PromptEncoder._embed_boxes3  sN  € à˜‘ˆØ�”Ð3˜Uœ[¨¨!¨œ_Ð3¨aÐ3°Ð3Ð3Ð3ˆå”Ô$×(Ò(¨°ÀJÐVWÐ(ÑXÔXˆØ×0Ò0°¸$Ô:OÐQUÔQfÐ9gÑhÔhÐØ˜˜˜˜A˜A˜A˜q ! ! !˜Ð$Ð$Ô$¨Ô(8Ô(?ÀÔ(BÑBÐ$Ð$Ñ$Ø˜˜˜˜A˜A˜A˜q ! ! !˜Ð$Ð$Ô$¨Ô(8Ô(?ÀÔ(BÑBÐ$Ð$Ñ$Ø'+Ô'=Ô'D×'NÒ'NÐO_Ð`aÐ`aÐ`aÐcdÐcdÐcdÐfgÐijÐijÐijÐ`jÔOkÑ'lÔ'lÐ˜˜˜˜A˜A˜A˜q ! ! !˜Ñ$ØÐr.   Úinput_pointsNÚinput_labelsÚinput_boxesÚinput_masksc                 óØ  — d}d}|�:|j         d         }|€t          d¦  «        ‚|                      |||du ¬¦  «        }|}|�?|j         d         }|                      |¦  «        }|€|}nt	          j        ||gd¬¦  «        }|�|                      |¦  «        }	nN| j        j         	                    dddd¦  «         
                    |d| j        d         | j        d         ¦  «        }	||	fS )	au  
        Embeds different types of prompts, returning both sparse and dense embeddings.

        Args:
            points (`torch.Tensor`, *optional*):
                point coordinates and labels to embed.
            boxes (`torch.Tensor`, *optional*):
                boxes to embed
            masks (`torch.Tensor`, *optional*):
                masks to embed
        Nr   r   z5If points are provided, labels must also be provided.)rò   r   rr   rn   )rN   rd   r£  r©  r*   r}   r‹  r�  rP   rØ   ru   r�  )
rK   rª  r«  r¬  r­  Úsparse_embeddingsr~   Úpoint_embeddingsÚbox_embeddingsr†  s
             r/   rX   zSam2PromptEncoder.forward?  s#  € ð$ !ÐØˆ
ØÐ#Ø%Ô+¨AÔ.ˆJØÐ#Ý Ð!XÑYÔYÐYØ#×1Ò1°,ÀÐS^ÐbfÐSfÐ1ÑhÔhÐØ 0ÐØÐ"Ø$Ô*¨1Ô-ˆJØ!×.Ò.¨{Ñ;Ô;ˆNØ Ð(Ø$2Ð!Ð!å$)¤IÐ/@À.Ð.QÐWXÐ$YÑ$YÔ$YÐ!ØÐ"Ø#Ÿš¨{Ñ;Ô;ÐÐà#Ô1Ô8×@Ò@ÀÀBÈÈ1ÑMÔM×TÒTØ˜B Ô 9¸!Ô <¸dÔ>WÐXYÔ>Zñ ô  Ðð !Ð"2Ð2Ð2r.   )r&   r'   r(   r   rB   r*   r   r‹   r£  r©  r9   rX   rY   rZ   s   @r/   rˆ  rˆ    s$  ø€ € € € € ðEÐ6ð Eð Eð Eð Eð Eð Eð E¤Lð ¸%¼,ð ÈTð ÐV[ÔVbð ð ð ð ð2
  %¤,ð 
 °5´<ð 
 ð 
 ð 
 ð 
 ð(3à˜EœL¨%¬,Ð6Ô7¸$Ñ>ð(3ð ”l TÑ)ð(3ð ”\ DÑ(ð	(3ð
 ”\ DÑ(ð(3ð 
ˆuŒ|˜Uœ\Ð)Ô	*ð(3ð (3ð (3ð (3ð (3ð (3ð (3ð (3r.   rˆ  c                   ó¦   ‡ — e Zd ZdZdˆ fd„	Z	 ddej        dej        dej        dej        dz  dee         d	e	ej        ej        f         fd
„Z
ˆ xZS )ÚSam2Attentionz‰
    SAM2's attention layer that allows for downscaling the size of the embedding after projection to queries, keys, and
    values.
    Nc                 ó.  •— t          ¦   «                              ¦   «          |€|j        n|}|| _        |j        | _        |j        |z  | _        |j        | _        | j        |j        z  | _        | j        dz  | _        d| _	        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        d S )NrÐ   F)rA   rB   Úattention_downsample_rater<   rD   Úinternal_dimrÎ   rÕ   r·   rÔ   rE   rÑ   Úq_projÚk_projÚv_projÚo_proj)rK   r<   Údownsample_raterL   s      €r/   rB   zSam2Attention.__init__p  sè   ø€ Ý‰Œ×ÒÑÔÐØ>MÐ>U˜&Ô:Ð:Ð[jˆØˆŒØ!Ô-ˆÔØ"Ô.°/ÑAˆÔØ#)Ô#=ˆÔ ØÔ)¨VÔ-GÑGˆŒØ”} dÑ*ˆŒØˆŒå”i Ô 0°$Ô2CÑDÔDˆŒÝ”i Ô 0°$Ô2CÑDÔDˆŒÝ”i Ô 0°$Ô2CÑDÔDˆŒÝ”i Ô 1°4Ô3CÑDÔDˆŒˆˆr.   r³   r´   rµ   Úattention_similarityrÁ   rj   c                 ó  — |j         d d…         \  }}||z  d| j        | j        f} |                      |¦  «        j        |Ž                      dd¦  «        } |                      |¦  «        j        |Ž                      dd¦  «        } |                      |¦  «        j        |Ž                      dd¦  «        }t          j	        | j
        j        t          ¦  «        }	t          | j
        ¦  «        r)|�'t          d         }	t                               d¦  «          |	| |||f|d| j        | j        dœ|¤Ž\  }
}|
                     ||d| j        | j        z  ¦  «                             ¦   «         }
|                      |
¦  «        }
|
|fS )Nr   rn   r   ÚsdpaznFalling back to SDPA for target-guided attention because Flash Attention does not support additive bias masks.r±   )r¶   r¸   r·   rÔ   )rN   rÎ   rÕ   r·  ró   r½   r¸  r¹  r   rÚ   r<   rÛ   rÄ   r   ÚloggerÚwarning_oncer·   rÔ   rØ   rÀ   rº  )rK   r³   r´   rµ   r¼  rÁ   r~   Úpoint_batch_sizeÚ	new_shaperÜ   rÃ   rÂ   s               r/   rX   zSam2Attention.forward€  s¸  € ð (-¤{°2°A°2¤Ñ$ˆ
Ð$ØÐ"2Ñ2°B¸Ô8PÐRVÔR_Ð`ˆ	à'�—’˜EÑ"Ô"Ô'¨Ð3×=Ò=¸aÀÑCÔCˆØ#ˆd�kŠk˜#ÑÔÔ# YÐ/×9Ò9¸!¸QÑ?Ô?ˆØ'�—’˜EÑ"Ô"Ô'¨Ð3×=Ò=¸aÀÑCÔCˆå(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðõ (¨¬Ñ4Ô4ð 	Ð9MÐ9Yõ #:¸&Ô"AÐÝ×ÒðHñô ð ð
 %8Ð$7ØØØØð	
%
ð
 0ØØ”LØ”nð
%
ð 
%
ð ð
%
ð 
%
Ñ!ˆ�\ð "×)Ò)ØÐ(¨"¨dÔ.FÈÌÑ.Vñ
ô 
ç
Š*‰,Œ,ð 	ð —k’k +Ñ.Ô.ˆà˜LÐ(Ð(r.   r‰   )r&   r'   r(   r)   rB   r*   r   r   r   r9   rX   rY   rZ   s   @r/   r³  r³  j  sÆ   ø€ € € € € ðð ð
Eð Eð Eð Eð Eð Eð* 59ð.)ð .)àŒ|ð.)ð Œ\ð.)ð Œ|ð	.)ð
 $œl¨TÑ1ð.)ð Ð+Ô,ð.)ð 
ˆuŒ|˜Uœ\Ð)Ô	*ð.)ð .)ð .)ð .)ð .)ð .)ð .)ð .)r.   r³  c                   óV   ‡ — e Zd Zddedefˆ fd„Zdedededed	ed
ee         fd„Z	ˆ xZ
S )ÚSam2TwoWayAttentionBlockFr<   Úskip_first_layer_pec                 ó  •— t          ¦   «                              ¦   «          t          |d¬¦  «        | _        t	          j        |j        ¦  «        | _        t          |¦  «        | _        t	          j        |j        ¦  «        | _	        t          |j        |j        |j        |j        ¬¦  «        | _        t	          j        |j        ¦  «        | _        t	          j        |j        ¦  «        | _        t          |¦  «        | _        || _        dS )aÛ  
        A transformer block with four layers:
            (1) self-attention of sparse inputs (2) cross attention of sparse inputs -> dense inputs (3) mlp block on
            sparse inputs (4) cross attention of dense inputs -> sparse inputs

        Arguments:
            config (`Sam2MaskDecoderConfig`):
                The configuration file used to instantiate the block
            attention_downsample_rate (*optionalk*, int, defaults to 2):
                The downsample ratio of the block used to reduce the inner dim of the attention.
            skip_first_layer_pe (*optional*, bool, defaults to `False`):
                Whether or not to skip the addition of the query_point_embedding on the first layer.
        r   )r»  )rã   N)rA   rB   r³  Ú	self_attnrE   r	  rD   r  Úcross_attn_token_to_imager  rÞ   Úmlp_dimÚnum_hidden_layersr  Úlayer_norm3Úlayer_norm4Úcross_attn_image_to_tokenrÅ  )rK   r<   rÅ  rL   s      €r/   rB   z!Sam2TwoWayAttentionBlock.__init__²  sÝ   ø€ õ 	‰Œ×ÒÑÔÐÝ& v¸qÐAÑAÔAˆŒÝœ<¨Ô(:Ñ;Ô;ˆÔå)6°vÑ)>Ô)>ˆÔ&Ýœ<¨Ô(:Ñ;Ô;ˆÔå"ØÔ ¤°Ô0BÈvÔOgð
ñ 
ô 
ˆŒõ œ<¨Ô(:Ñ;Ô;ˆÔåœ<¨Ô(:Ñ;Ô;ˆÔÝ)6°vÑ)>Ô)>ˆÔ&à#6ˆÔ Ð Ð r.   ÚqueriesÚkeysÚquery_point_embeddingÚkey_point_embeddingr¼  rÁ   c                 ó"  — | j         r|                      |||¬¦  «        \  }}n%||z   }|                      |||¬¦  «        \  }	}||	z   }|                      |¦  «        }||z   }||z   }
|                      ||
||¬¦  «        \  }	}||	z   }|                      |¦  «        }|                      |¦  «        }||z   }|                      |¦  «        }||z   }||z   }
|                      |
||¬¦  «        \  }	}||	z   }|                      |¦  «        }|||	fS )N©r³   r´   rµ   )r³   r´   rµ   r¼  )	rÅ  rÇ  r  rÈ  r  r  rË  rÍ  rÌ  )rK   rÎ  rÏ  rÐ  rÑ  r¼  rÁ   rT   r³   Úattn_outr´   Úmlp_outs               r/   rX   z Sam2TwoWayAttentionBlock.forwardÑ  s^  € ð Ô#ð 	)ØŸš¨g¸7È'˜ÑRÔR‰JˆG�Q�QàÐ3Ñ3ˆEØŸ.š.¨u¸%Àw˜.ÑOÔO‰KˆH�aØ Ñ(ˆGØ×"Ò" 7Ñ+Ô+ˆð Ð/Ñ/ˆØÐ(Ñ(ˆà×4Ò4Ø˜S¨ÐCWð 5ñ 
ô 
‰ˆ�!ð ˜HÑ$ˆà×"Ò" 7Ñ+Ô+ˆð —(’(˜7Ñ#Ô#ˆØ˜GÑ#ˆØ×"Ò" 7Ñ+Ô+ˆð Ð/Ñ/ˆØÐ(Ñ(ˆà×4Ò4¸3ÀEÐQXÐ4ÑYÔY‰ˆ�!Ø�h‰ˆà×Ò Ñ%Ô%ˆØ˜˜hÐ&Ð&r.   )F)r&   r'   r(   r   r‹   rB   r   r   r   rX   rY   rZ   s   @r/   rÄ  rÄ  ±  s¥   ø€ € € € € ð7ð 7Ð4ð 7È4ð 7ð 7ð 7ð 7ð 7ð 7ð>*'àð*'ð ð*'ð  &ð	*'ð
 $ð*'ð %ð*'ð Ð+Ô,ð*'ð *'ð *'ð *'ð *'ð *'ð *'ð *'r.   rÄ  c                   óZ   ‡ — e Zd Zdefˆ fd„Z	 ddededededee         d	ee	z  fd
„Z
ˆ xZS )ÚSam2TwoWayTransformerr<   c                 óŠ  •— t          ¦   «                              ¦   «          || _        |j        | _        t	          j        ¦   «         | _        t          | j        ¦  «        D ]/}| j                             t          ||dk    ¬¦  «        ¦  «         Œ0t          |¦  «        | _        t	          j        |j        ¦  «        | _        d S )Nr   )rÅ  )rA   rB   r<   rÊ  rE   r—   rì   r§   rš   rÄ  r³  Úfinal_attn_token_to_imager	  rD   Úlayer_norm_final_attn)rK   r<   r¬   rL   s      €r/   rB   zSam2TwoWayTransformer.__init__ÿ  s­   ø€ Ý‰Œ×ÒÑÔÐØˆŒà!'Ô!9ˆÔÝ”m‘o”oˆŒå�tÔ-Ñ.Ô.ð 	_ð 	_ˆAØŒK×ÒÕ7¸ÐUVÐZ[ÒU[Ð]Ñ]Ô]Ñ^Ô^Ð^Ð^å)6°vÑ)>Ô)>ˆÔ&Ý%'¤\°&Ô2DÑ%EÔ%EˆÔ"Ð"Ð"r.   Nr°  r5   Úimage_positional_embeddingsr¼  rÁ   rj   c           
      óè  — |€t          d¦  «        ‚|                     d¦  «                             dd¦  «                             d¦  «        }|                     d¦  «                             dd¦  «                             d¦  «        }|}|}| j        D ]}	|�||z  } |	d|||||dœ|¤Ž\  }}}
Œ||z   }||z   }|                      |||¬¦  «        \  }}
||z   }|                      |¦  «        }||fS )Nz&You have to specify an image_embeddingr   r   )rÎ  rÏ  rÐ  rÑ  r¼  rÓ  r-   )rd   r|   r½   r¡  rì   rÙ  rÚ  )rK   r°  r5   rÛ  r¼  Útarget_embeddingrÁ   rÎ  rÏ  rï   rT   r³   r´   rÔ  s                 r/   rX   zSam2TwoWayTransformer.forward  sL  € ð Ð#ÝÐEÑFÔFÐFà+×3Ò3°AÑ6Ô6×@Ò@ÀÀAÑFÔF×PÒPÐQRÑSÔSÐØ&A×&IÒ&IÈ!Ñ&LÔ&L×&VÒ&VÐWXÐZ[Ñ&\Ô&\×&fÒ&fÐghÑ&iÔ&iÐ#ð #ˆØˆð ”[ð 	ð 	ˆEØÐ+ØÐ+Ñ+�à$˜uð  ØØØ&6Ø$?Ø%9ð ð  ð ð ð  ÑˆG�T˜1˜1ð Ð*Ñ*ˆØÐ0Ñ0ˆà×4Ò4¸5ÀcÐQUÐ4ÑVÔV‰ˆ�!à˜HÑ$ˆØ×,Ò,¨WÑ5Ô5ˆØ˜ˆ}Ðr.   r‰   )r&   r'   r(   r   rB   r   r   r   r9   r
   rX   rY   rZ   s   @r/   r×  r×  þ  s¯   ø€ € € € € ðFÐ4ð Fð Fð Fð Fð Fð Fð& ð(ð (à ð(ð !ð(ð &,ð	(ð
 %ð(ð Ð+Ô,ð(ð 
�Ñ	 ð(ð (ð (ð (ð (ð (ð (ð (r.   r×  c                   óR   ‡ — e Zd ZdZdddœˆ fd„
Zdej        dej        fˆ fd„Zˆ xZS )	rƒ  aA  LayerNorm that supports two data formats: channels_last (default) or channels_first.
    The ordering of the dimensions in the inputs. channels_last corresponds to inputs with shape (batch_size, height,
    width, channels) while channels_first corresponds to inputs with shape (batch_size, channels, height, width).
    rm   Úchannels_lastr}  c                óz   •—  t          ¦   «         j        |fd|i|¤Ž |dvrt          d|› �¦  «        ‚|| _        d S )Nr‚   )rß  r|  zUnsupported data format: )rA   rB   ÚNotImplementedErrorr~  )rK   Únormalized_shaper‚   r~  rÁ   rL   s        €r/   rB   zSam2LayerNorm.__init__=  sY   ø€ Ø�‰ŒÔÐ)Ð=Ð=¨sÐ=°fÐ=Ð=Ð=ØÐAÐAÐAÝ%Ð&OÀ+Ð&OÐ&OÑPÔPÐPØ&ˆÔÐÐr.   Úfeaturesrj   c                 ó  •— | j         dk    rR|                     dddd¦  «        }t          ¦   «                              |¦  «        }|                     dddd¦  «        }n!t          ¦   «                              |¦  «        }|S )zŒ
        Args:
            features: Tensor of shape (batch_size, channels, height, width) OR (batch_size, height, width, channels)
        r|  r   r   r   r   )r~  rR   rA   rX   )rK   rã  rL   s     €r/   rX   zSam2LayerNorm.forwardC  sw   ø€ ð
 ÔÐ/Ò/Ð/Ø×'Ò'¨¨1¨a°Ñ3Ô3ˆHÝ‘w”w—’ xÑ0Ô0ˆHØ×'Ò'¨¨1¨a°Ñ3Ô3ˆHˆHå‘w”w—’ xÑ0Ô0ˆHØˆr.   )	r&   r'   r(   r)   rB   r*   r   rX   rY   rZ   s   @r/   rƒ  rƒ  7  sƒ   ø€ € € € € ðð ð
 15À/ð 'ð 'ð 'ð 'ð 'ð 'ð 'ð ¤ð °´ð ð ð ð ð ð ð ð ð ð r.   rƒ  c                   ó  ‡ — e Zd Zdefˆ fd„Z	 	 ddej        dej        dej        dej        ded	eej                 d
ej        dz  dej        dz  de	e
         deej        ej        ej        ej        f         fd„Zd„ Zd„ Zˆ xZS )ÚSam2MaskDecoderr<   c                 óê  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        |j        dz   | _        t          j        d| j        ¦  «        | _        t          j        | j        | j        ¦  «        | _	        t          |¦  «        | _        t          j        | j        | j        dz  dd¬¦  «        | _        t          j        | j        dz  | j        dz  dd¬¦  «        | _        t          | j        dz  d¬¦  «        | _        t          j        ¦   «         | _        g }t'          | j        ¦  «        D ]*}|t)          | j        | j        | j        dz  d¦  «        gz  }Œ+t          j        |¦  «        | _        t)          | j        |j        | j        |j        d	¬
¦  «        | _        t          j        |j        |j        dz  dd¬¦  «        | _        t          j        |j        |j        dz  dd¬¦  «        | _        t          j        d| j        ¦  «        | _        t)          | j        | j        dd¦  «        | _        |j        | _        |j         | _         |j!        | _!        d S )Nr   rq   r   r{  é   r|  )r~  r   T)rå   )"rA   rB   r<   rD   Únum_multimask_outputsÚnum_mask_tokensrE   rŒ  Ú	iou_tokenÚmask_tokensr×  ÚtransformerÚConvTranspose2dÚupscale_conv1Úupscale_conv2rƒ  Úupscale_layer_normÚGELUrä   r§   rÞ   r—   Úoutput_hypernetworks_mlpsÚiou_head_hidden_dimÚiou_head_depthÚiou_prediction_headrF   Úconv_s0Úconv_s1Úobj_score_tokenÚpred_obj_score_headÚdynamic_multimask_via_stabilityÚ!dynamic_multimask_stability_deltaÚ"dynamic_multimask_stability_thresh)rK   r<   Ú	mlps_listrT   rL   s       €r/   rB   zSam2MaskDecoder.__init__R  sJ  ø€ Ý‰Œ×ÒÑÔÐØˆŒØ!Ô-ˆÔà%+Ô%AˆÔ"Ø%Ô;¸aÑ?ˆÔåœ a¨Ô)9Ñ:Ô:ˆŒÝœ<¨Ô(<¸dÔ>NÑOÔOˆÔå0°Ñ8Ô8ˆÔõ  Ô/°Ô0@À$ÔBRÐVWÑBWÐefÐopÐqÑqÔqˆÔÝÔ/°Ô0@ÀAÑ0EÀtÔGWÐ[\ÑG\ÐjkÐtuÐvÑvÔvˆÔÝ"/°Ô0@ÀAÑ0EÐScÐ"dÑ"dÔ"dˆÔÝœ'™)œ)ˆŒàˆ	Ý�tÔ+Ñ,Ô,ð 	ið 	iˆAØ�/¨$Ô*:¸DÔ<LÈdÔN^ÐbcÑNcÐefÑgÔgÐhÑhˆIˆIÝ)+¬°yÑ)AÔ)AˆÔ&Ý#2ØÔØÔ&ØÔ ØÔ!Øð$
ñ $
ô $
ˆÔ õ ”y Ô!3°VÔ5GÈ1Ñ5LÐZ[ÐdeÐfÑfÔfˆŒÝ”y Ô!3°VÔ5GÈ1Ñ5LÐZ[ÐdeÐfÑfÔfˆŒå!œ|¨A¨tÔ/?Ñ@Ô@ˆÔÝ#2°4Ô3CÀTÔEUÐWXÐZ[Ñ#\Ô#\ˆÔ à/5Ô/UˆÔ,Ø17Ô1YˆÔ.Ø28Ô2[ˆÔ/Ð/Ð/r.   Nr5   rÛ  Úsparse_prompt_embeddingsÚdense_prompt_embeddingsÚmultimask_outputÚhigh_resolution_featuresr¼  rÝ  rÁ   rj   c	           
      ó  — |j         \  }
}}}|j         d         }t          j        | j        j        | j        j        | j        j        gd¬¦  «        }|                     |
|dd¦  «        }|j         d         dk    rt          j        ||fd¬¦  «        }n|}|                     | j        j        j	        ¦  «        }||z   }| 
                    |d¬¦  «        }| 
                    |d¦  «        } | j        d	|||||dœ|	¤Ž\  }}|dd…dd…ddd…f         }|dd…dd…dd| j        z   …dd…f         }|                     dd¦  «                             |
|z  |||¦  «        }|\  }}| 
                    |d¬¦  «        }| 
                    |d¬¦  «        }|                      |¦  «        |z   }|                      |                      |¦  «        ¦  «        }|                      |                      |¦  «        |z   ¦  «        }g }t'          | j        ¦  «        D ].}| j        |         }| ||dd…dd…|dd…f         ¦  «        gz  }Œ/t          j        |d¬¦  «        }|j         \  }}}}|                     |
||||z  ¦  «        }||z                       |
|d||¦  «        }|                      |¦  «        }|                      |dd…dd…ddd…f         ¦  «        }|r5t1          dd¦  «        }|dd…dd…|dd…dd…f         }|dd…dd…|f         }nl| j        r1| j        s*t1          dd¦  «        }|                      ||¦  «        \  }}n4t1          dd¦  «        }|dd…dd…|dd…dd…f         }|dd…dd…|f         }|dd…dd…|f         } ||| |fS )
aÛ  
        Predict masks given image and prompt embeddings.

        Args:
            image_embeddings (`torch.Tensor`):
                The embeddings from the image encoder.
            image_positional_embeddings (`torch.Tensor`):
                Positional encoding with the shape of image_embeddings.
            sparse_prompt_embeddings (`torch.Tensor`):
                The embeddings of the points and boxes.
            dense_prompt_embeddings (`torch.Tensor`):
                The embeddings of the mask inputs.
            multimask_output (`bool`):
                Whether to return multiple masks or a single mask.
            high_resolution_features (`list[torch.Tensor]`, *optional*):
                The high-resolution features from the vision encoder.
            attention_similarity (`torch.Tensor`, *optional*):
                The attention similarity tensor.
            target_embedding (`torch.Tensor`, *optional*):
                The target embedding.
        r   r   rr   r   )r°  r5   rÛ  r¼  rÝ  Nr   rn   r-   )rN   r*   r}   rù  rP   rë  rì  ÚrepeatrO   rQ   Úrepeat_interleaverí  rê  r½   ró   rï  rä   rñ  rð  r§   ró  ry   rö  rú  Úslicerû  r»   Ú _dynamic_multimask_via_stability)!rK   r5   rÛ  rÿ  r   r  r  r¼  rÝ  rÁ   r~   rC   rU   rV   rÁ  Úoutput_tokensÚtokensr°  Úiou_token_outÚmask_tokens_outÚfeat_s0Úfeat_s1Úupscaled_embeddingÚhyper_in_listr¬   Úcurrent_mlpÚhyper_inrT   r…  Úiou_predr4   Ú
mask_sliceÚsam_tokens_outs!                                    r/   rX   zSam2MaskDecoder.forward{  sˆ  € ðB 3CÔ2HÑ/ˆ
�L &¨%Ø3Ô9¸!Ô<Ðåœ	àÔ$Ô+Ø”Ô%ØÔ Ô'ðð
 ð
ñ 
ô 
ˆð &×,Ò,¨ZÐ9IÈ1ÈaÑPÔPˆà#Ô)¨!Ô,°Ò1Ð1Ý”Y Ð/GÐHÈaÐPÑPÔPˆFˆFà"ˆFØ!Ÿ9š9 T¤^Ô%:Ô%@ÑAÔAÐð ,Ð.EÑEÐØ+×=Ò=Ð>NÐTUÐ=ÑVÔVÐØ&A×&SÒ&SÐTdÐfgÑ&hÔ&hÐ#à-=¨TÔ-=ð .
Ø-Ø-Ø(CØ!5Ø-ð.
ð .
ð ð.
ð .
Ñ*ÐÐ*ð )¨¨¨¨A¨A¨A¨q°!°!°!¨Ô4ˆØ*¨1¨1¨1¨a¨a¨a°°a¸$Ô:NÑ6NÐ1OÐQRÐQRÐQRÐ+RÔSˆð ,×5Ò5°a¸Ñ;Ô;×@Ò@ØÐ)Ñ)¨<¸Àñ
ô 
Ðð 4Ñˆ�Ø×+Ò+Ð,<À!Ð+ÑDÔDˆØ×+Ò+Ð,<À!Ð+ÑDÔDˆØ!×/Ò/Ð0@ÑAÔAÀGÑKÐØ!Ÿ_š_¨T×-DÒ-DÐEWÑ-XÔ-XÑYÔYÐØ!Ÿ_š_¨T×-?Ò-?Ð@RÑ-SÔ-SÐV]Ñ-]Ñ^Ô^Ðà,.ˆÝ�tÔ+Ñ,Ô,ð 	Hð 	HˆAØÔ8¸Ô;ˆKØ˜k˜k¨/¸!¸!¸!¸Q¸Q¸QÀÀ1À1À1¸*Ô*EÑFÔFÐGÑGˆMˆMÝ”;˜}°!Ð4Ñ4Ô4ˆà);Ô)AÑ&ˆˆ<˜ Ø/×4Ò4°ZÐAQÐS_ÐagÐjoÑaoÑpÔpÐØÐ.Ñ.×4Ò4°ZÐAQÐSUÐW]Ð_dÑeÔeˆð ×+Ò+¨MÑ:Ô:ˆØ"×6Ò6Ð7GÈÈÈÈ1È1È1ÈaÐQRÐQRÐQRÈ
Ô7SÑTÔTÐð ð 
	2Ý˜q $™œˆJØ˜!˜!˜!˜Q˜Q˜Q 
¨A¨A¨A¨q¨q¨qÐ0Ô1ˆEØ    1 1 1 jÐ 0Ô1ˆHˆHØÔ1ð 	2¸$¼-ð 	2Ý˜q !™œˆJØ"×CÒCÀEÈ8ÑTÔT‰OˆE�8�8å˜q !™œˆJØ˜!˜!˜!˜Q˜Q˜Q 
¨A¨A¨A¨q¨q¨qÐ0Ô1ˆEØ    1 1 1 jÐ 0Ô1ˆHà(¨¨¨¨A¨A¨A¨zÐ)9Ô:ˆà�h Ð0CÐCÐCr.   c                 ó*  — |                      d¦  «        }| j        }t          j        ||k    d¬¦  «                             ¦   «         }t          j        || k    d¬¦  «                             ¦   «         }t          j        |dk    ||z  d¦  «        }|S )zz
        Compute stability scores of the mask logits based on the IoU between upper and
        lower thresholds.
        r×   rn   rr   r   g      ð?)r|   rü  r*   ÚsumrŒ   rž  )rK   Úmask_logitsÚstability_deltaÚarea_iÚarea_uÚstability_scoress         r/   Ú_get_stability_scoresz%Sam2MaskDecoder._get_stability_scoresê  sŽ   € ð
 "×)Ò)¨"Ñ-Ô-ˆØÔ@ˆÝ”˜;¨Ò8¸bÐAÑAÔA×GÒGÑIÔIˆÝ”˜;¨/Ð)9Ò9¸rÐBÑBÔB×HÒHÑJÔJˆÝ œ; v°¢z°6¸F±?ÀCÑHÔHÐØÐr.   c           	      ó8  — |dd…dd…dd…dd…dd…f         }|dd…dd…dd…f         }t          j        |d¬¦  «        }|                     d¦  «                             d¦  «                             d¦  «        }|                     ddd|                     d¦  «        |                     d¦  «        ¦  «        }t          j        |d|¦  «        }t          j        |d|                     d¦  «        ¦  «        }|dd…dd…dd…dd…dd…f         }	|dd…dd…dd…f         }
|                      |	¦  «        }|| j        k    }t          j        |d          	                    |	¦  «        |	|¦  «        }t          j        | 	                    |
¦  «        |
|¦  «        }||fS )	as  
        When outputting a single mask, if the stability score from the current single-mask
        output (based on output token 0) falls below a threshold, we instead select from
        multi-mask outputs (based on output token 1~3) the mask with the highest predicted
        IoU score. This is intended to ensure a valid mask for both clicking and tracking.
        Nr   rn   rr   r×   r   r   ).NN)
r*   Úargmaxr¡  ru   rU  Úgatherr  rý  rž  r¦  )rK   Úall_mask_logitsÚall_iou_scoresÚmultimask_logitsÚmultimask_iou_scoresÚbest_scores_indsÚbest_scores_inds_expandedÚbest_multimask_logitsÚbest_multimask_iou_scoresÚsinglemask_logitsÚsinglemask_iou_scoresr  Ú	is_stableÚmask_logits_outÚiou_scores_outs                  r/   r  z0Sam2MaskDecoder._dynamic_multimask_via_stabilityö  sß  € ð +¨1¨1¨1¨a¨a¨a°°°°Q°Q°Q¸¸¸¨>Ô:ÐØ-¨a¨a¨a°°°°A°B°B¨hÔ7ÐÝ œ<Ð(<À"ÐEÑEÔEÐØ$4×$>Ò$>¸rÑ$BÔ$B×$LÒ$LÈRÑ$PÔ$P×$ZÒ$ZÐ[]Ñ$^Ô$^Ð!Ø$=×$DÒ$DØ��AÐ'×,Ò,¨RÑ0Ô0Ð2B×2GÒ2GÈÑ2KÔ2Kñ%
ô %
Ð!õ !&¤Ð-=¸qÐB[Ñ \Ô \ÐÝ$)¤LÐ1EÀqÐJZ×JdÒJdÐegÑJhÔJhÑ$iÔ$iÐ!ð ,¨A¨A¨A¨q¨q¨q°!°A°#°q°q°q¸!¸!¸!¨OÔ<ÐØ .¨q¨q¨q°!°!°!°Q°q°S¨yÔ 9ÐØ×5Ò5Ð6GÑHÔHÐØ$¨Ô(OÒOˆ	õ  œ+Ø�oÔ&×0Ò0Ð1BÑCÔCØØ!ñ
ô 
ˆõ
 œØ×ÒÐ 5Ñ6Ô6Ø!Ø%ñ
ô 
ˆð
  Ð.Ð.r.   )NN)r&   r'   r(   r   rB   r*   r   r‹   Úlistr   r   r9   rX   r  r  rY   rZ   s   @r/   ræ  ræ  Q  sF  ø€ € € € € ð'\Ð4ð '\ð '\ð '\ð '\ð '\ð '\ðb 59Ø04ðmDð mDàœ,ðmDð &+¤\ðmDð #(¤,ð	mDð
 "'¤ðmDð ðmDð #' u¤|Ô"4ðmDð $œl¨TÑ1ðmDð  œ,¨Ñ-ðmDð Ð+Ô,ðmDð 
ˆuŒ|˜Uœ\¨5¬<¸¼ÐEÔ	FðmDð mDð mDð mDð^
 ð 
 ð 
 ð#/ð #/ð #/ð #/ð #/ð #/ð #/r.   ræ  z”
    Segment Anything Model 2 (SAM 2) for generating segmentation masks, given an input image and
    input points and labels, boxes, or masks.
    c                   ó°  ‡ — e Zd ZdZd eed¬¦  «        iZi Zdefˆ fd„Z	d„ Z
dej        fd	„Z ej        ¦   «         d
ej        dee         deej                 fd„¦   «         Z ej        ¦   «         	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  fd„¦   «         Zeee	 	 	 	 	 	 	 	 	 dd
ej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dedej        dz  dej        dz  dee         defd„¦   «         ¦   «         ¦   «         Zeed
ej        dee         deez  fd„¦   «         ¦   «         Zˆ xZ S )r5  )r%  Útextr8   r   )Úindexr<   c                 óF  •— t          ¦   «                              |¦  «         t          |j        ¦  «        | _        t          j        |j        ¦  «        | _        t          |j        ¦  «        | _
        |j        |j        _        t          |j        ¦  «        | _        |j        j        | _        |j        j        | _        |j        j        | _        t&          j                             t'          j        dd| j        ¦  «        ¦  «        | _        |                      ¦   «          d S )Nr   )rA   rB   r2  Úprompt_encoder_configÚshared_image_embeddingr   re  Úvision_configÚvision_encoderrˆ  Úprompt_encoderrÛ   Úmask_decoder_configræ  Úmask_decoderri  Úbackbone_feature_sizesr•   rá   r*   rE   rD  rE  r6  rM  rj  s     €r/   rB   zSam2Model.__init__'  sß   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý&=¸fÔ>ZÑ&[Ô&[ˆÔ#Ý'Ô3°FÔ4HÑIÔIˆÔÝ/°Ô0LÑMÔMˆÔà:@Ô:UˆÔ"Ô7Ý+¨FÔ,FÑGÔGˆÔà"(Ô"6Ô"IˆÔØ&,Ô&:Ô&QˆÔ#à Ô.Ô>ˆŒÝ#(¤8×#5Ò#5µe´kÀ!ÀQÈÌÑ6XÔ6XÑ#YÔ#YˆÔ à�ŠÑÔÐÐÐr.   c                 ó4   — | j                              ¦   «         S r‰   )r5  rQ  rP  s    r/   rQ  zSam2Model.get_input_embeddings8  s   € ØÔ"×7Ò7Ñ9Ô9Ð9r.   rj   c                 óÆ  — | j         j        }| j        j        j        }| j        j        j        }t          j        |||¬¦  «        }|                     d¬¦  «        dz
  }|                     d¬¦  «        dz
  }||d         z  }||d         z  }|                      t          j	        ||gd¬¦  «        ¦  «        }| 
                    ddd¦  «                             d¦  «        S )N)rh   rQ   r   rr   r™  r   rn   r   )r6  r�  r3  r4  rh   rQ   r*   Úonesrv   ry   rR   r¡  )rK   rU  Útarget_deviceÚtarget_dtypeÚgridr   r€   r4  s           r/   Ú$get_image_wide_positional_embeddingsz.Sam2Model.get_image_wide_positional_embeddings;  sÛ   € ØÔ"Ô7ˆØÔ3ÔHÔOˆØÔ2ÔGÔMˆÝŒz˜$ }¸LÐIÑIÔIˆØ—+’+ !�+Ñ$Ô$ sÑ*ˆØ—+’+ !�+Ñ$Ô$ sÑ*ˆØ˜D œGÑ#ˆØ˜D œGÑ#ˆà#×:Ò:½5¼;ÈÐQXÐGYÐ_aÐ;bÑ;bÔ;bÑcÔcÐØ#×+Ò+¨A¨q°!Ñ4Ô4×>Ò>¸qÑAÔAÐAr.   rS   rÁ   c                 ó¸   ‡— |j         d         Š | j        |fddi|¤Ž}|j        }|d         | j        z   |d<   ˆfd„t	          || j        ¦  «        D ¦   «         }|S )zý
        Returns the image embeddings by passing the pixel values through the vision encoder.

        Args:
            pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
                Input pixel values
        r   Úreturn_dictTrn   c                 ó^   •— g | ])\  }} |                      d dd¦  «        j        ‰dg|¢R Ž ‘Œ*S ©r   r   r   rn   ©rR   ró   ©rè   ÚfeatÚ	feat_sizer~   s      €r/   ré   z2Sam2Model.get_image_embeddings.<locals>.<listcomp>]  sT   ø€ ð 
ð 
ð 
á��ið 'ˆD�LŠL˜˜A˜qÑ!Ô!Ô& z°2ÐB¸	ÐBÐBÐBð
ð 
ð 
r.   )rN   Úget_image_featuresr$   r6  rY  r9  )rK   rS   rÁ   Úimage_outputsÚfeature_mapsr5   r~   s         @r/   Úget_image_embeddingszSam2Model.get_image_embeddingsH  s‘   ø€ ð "Ô'¨Ô*ˆ
Ø/˜Ô/°ÐYÐYÈ$ÐYÐRXÐYÐYˆØ$Ô6ˆð (¨Ô+¨dÔ.FÑFˆ�RÑð
ð 
ð 
ð 
å#& |°TÔ5PÑ#QÔ#Qð
ñ 
ô 
Ðð
  Ðr.   Nrª  r«  r¬  r­  c                 ó8   — |                       ||||¬¦  «        }|S )a  
        Returns the prompt embeddings by passing the input points, labels, boxes and masks through the prompt encoder.

        Args:
            input_points (`torch.FloatTensor` of shape `(batch_size, point_batch_size, num_points_per_image, 2)`):
                Optional input points for the prompt encoder. The padding of the point is automatically done by the
                processor. `point_batch_size` refers to the number of masks that we want the model to predict per
                point. The model will output `point_batch_size` times 3 masks in total.
            input_labels (`torch.LongTensor` of shape `(batch_size, point_batch_size, num_points_per_image)`):
                Optional input labels for the prompt encoder. The padding of the labels is automatically done by the
                processor, or can be fed by the user.
            input_boxes (`torch.FloatTensor` of shape `(batch_size, num_boxes_per_image, 4)`):
                Optional input boxes for the prompt encoder. The padding of the boxes is automatically done by the
                processor. users can also pass manually the input boxes.
            input_masks (`torch.LongTensor` of shape `(batch_size, image_size, image_size)`):
                Optional input masks for the prompt encoder.
        ©rª  r«  r¬  r­  )r6  )rK   rª  r«  r¬  r­  Úprompt_outputs         r/   Úget_prompt_embeddingszSam2Model.get_prompt_embeddingsd  s2   € ð2 ×+Ò+Ø%Ø%Ø#Ø#ð	 ,ñ 
ô 
ˆð Ðr.   Tr5   r  r¼  rÝ  c
                 óì  ‡— |du |du z  st          d¦  «        ‚|�J|�H|j        d         |j        d         k    r,t          d|j        d         › d|j        d         › d�¦  «        ‚|                      ¦   «         }|�|j        d         n|d         j        d         Š|                     ‰ddd¦  «        }d}d}|�Y | j        |fd	d
i|
¤Ž}|j        }|j        }|j        }|d         | j        z   |d<   ˆfd„t          || j
        ¦  «        D ¦   «         }|�8|€6t          j        |dd…dd…dd…df         t          j        |j        ¬¦  «        }|€a|€_t          j        ‰ddd|d         j        |d         j        ¬¦  «        }t          j        ‰ddt          j        |d         j        ¬¦  «         }|�j|j        dd…         | j        j        k    rMt+          j        |                     ¦   «         | j        j        ddd
¬¦  «                             |j        ¦  «        }|                      ||||¬¦  «        \  }} | j        d|d         |||||dd…         ||	dœ|
¤Ž\  }}}}t5          ||||||¬¦  «        S )aÜ  
        input_points (`torch.FloatTensor` of shape `(batch_size, num_points, 2)`):
            Input 2D spatial points, this is used by the prompt encoder to encode the prompt. Generally yields to much
            better results. The points can be obtained by passing a list of list of list to the processor that will
            create corresponding `torch` tensors of dimension 4. The first dimension is the image batch size, the
            second dimension is the point batch size (i.e. how many segmentation masks do we want the model to predict
            per input point), the third dimension is the number of points per segmentation mask (it is possible to pass
            multiple points for a single mask), and the last dimension is the x (vertical) and y (horizontal)
            coordinates of the point. If a different number of points is passed either for each image, or for each
            mask, the processor will create "PAD" points that will correspond to the (0, 0) coordinate, and the
            computation of the embedding will be skipped for these points using the labels.
        input_labels (`torch.LongTensor` of shape `(batch_size, point_batch_size, num_points)`):
            Input labels for the points, this is used by the prompt encoder to encode the prompt. According to the
            official implementation, there are 3 types of labels

            - `1`: the point is a point that contains the object of interest
            - `0`: the point is a point that does not contain the object of interest
            - `-1`: the point corresponds to the background

            We added the label:

            - `-10`: the point is a padding point, thus should be ignored by the prompt encoder

            The padding labels should be automatically done by the processor.
        input_boxes (`torch.FloatTensor` of shape `(batch_size, num_boxes, 4)`):
            Input boxes for the points, this is used by the prompt encoder to encode the prompt. Generally yields to
            much better generated masks. The boxes can be obtained by passing a list of list of list to the processor,
            that will generate a `torch` tensor, with each dimension corresponding respectively to the image batch
            size, the number of boxes per image and the coordinates of the top left and bottom right point of the box.
            In the order (`x1`, `y1`, `x2`, `y2`):

            - `x1`: the x coordinate of the top left point of the input box
            - `y1`: the y coordinate of the top left point of the input box
            - `x2`: the x coordinate of the bottom right point of the input box
            - `y2`: the y coordinate of the bottom right point of the input box
        input_masks (`torch.FloatTensor` of shape `(batch_size, image_size, image_size)`):
            SAM model also accepts segmentation masks as input. The mask will be embedded by the prompt encoder to
            generate a corresponding embedding, that will be fed later on to the mask decoder. These masks needs to be
            manually fed by the user, and they need to be of shape (`batch_size`, `image_size`, `image_size`).
        image_embeddings (`torch.FloatTensor` of shape `(batch_size, output_channels, window_size, window_size)`):
            Image embeddings, this is used by the mask decoder to generate masks and iou scores. For more memory
            efficient computation, users can first retrieve the image embeddings using the `get_image_embeddings`
            method, and then feed them to the `forward` method instead of feeding the `pixel_values`.
        multimask_output (`bool`, *optional*):
            In the original implementation and paper, the model always outputs 3 masks per image (or per point / per
            bounding box if relevant). However, it is possible to just output a single mask, that corresponds to the
            "best" mask, by specifying `multimask_output=False`.
        attention_similarity (`torch.FloatTensor`, *optional*):
            Attention similarity tensor, to be provided to the mask decoder for target-guided attention in case the
            model is used for personalization as introduced in [PerSAM](https://huggingface.co/papers/2305.03048).
        target_embedding (`torch.FloatTensor`, *optional*):
            Embedding of the target concept, to be provided to the mask decoder for target-semantic prompting in case
            the model is used for personalization as introduced in [PerSAM](https://huggingface.co/papers/2305.03048).

        Example:

        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoModel, AutoProcessor

        >>> model = AutoModel.from_pretrained("danelcsb/sam2.1_hiera_tiny")
        >>> processor = AutoProcessor.from_pretrained("danelcsb/sam2.1_hiera_tiny")

        >>> url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/transformers/model_doc/sam-car.png"
        >>> with httpx.stream("GET", url) as response:
        ...     raw_image = Image.open(BytesIO(response.read())).convert("RGB")
        >>> input_points = [[[400, 650]]]  # 2D location of a window on the car
        >>> inputs = processor(images=raw_image, input_points=input_points, return_tensors="pt")

        >>> # Get segmentation mask
        >>> outputs = model(**inputs)

        >>> # Postprocess masks
        >>> masks = processor.post_process_masks(
        ...     outputs.pred_masks, inputs["original_sizes"]
        ... )
        ```
        NzAExactly one of pixel_values or image_embeddings must be provided.r   zGYou should provide as many bounding boxes as input points per box. Got z and ú.r   rn   rB  Tc                 ó^   •— g | ])\  }} |                      d dd¦  «        j        ‰dg|¢R Ž ‘Œ*S rD  rE  rF  s      €r/   ré   z%Sam2Model.forward.<locals>.<listcomp>ÿ  sT   ø€ ð  ð  ð  á#�D˜)ð +�—’˜Q  1Ñ%Ô%Ô*¨:°rÐF¸IÐFÐFÐFð ð  ð  r.   rl   r   r×   FÚbilinear)rU  r¤   r£   r¥   rN  )r5   rÛ  rÿ  r   r  r  r¼  rÝ  )r2   r3   r4   r5   r6   r7   r-   )rd   rN   r@  r  rI  r$   rŸ   r!  r6  rY  r9  r*   Ú	ones_likerŠ   rh   rE  rQ   r<  Úint32r6  r‘  r¨   r©   rŒ   rO   r8  r1   )rK   rS   rª  r«  r¬  r­  r5   r  r¼  rÝ  rÁ   rÛ  r7   r6   rJ  rK  r¯  r†  Úlow_res_multimasksr2   rT   r4   r~   s                         @r/   rX   zSam2Model.forward…  s‚  ø€ ð@  Ð%Ð*:¸dÐ*BÑCð 	bÝÐ`ÑaÔaÐaØÐ#¨Ð(?ØÔ! !Ô$¨Ô(9¸!Ô(<Ò<Ð<Ý ð RÐ^jÔ^pÐqrÔ^sð  Rð  Rð  {Fô  {Lð  MNô  {Oð  Rð  Rð  Rñô ð ð '+×&OÒ&OÑ&QÔ&QÐ#à.:Ð.F�\Ô'¨Ô*Ð*ÐL\Ð]_ÔL`ÔLfÐghÔLiˆ
Ø&A×&HÒ&HÈÐUVÐXYÐ[\Ñ&]Ô&]Ð#à ÐØ#ÐàÐ#Ø5L°TÔ5LÈ\Ð5vÐ5vÐgkÐ5vÐouÐ5vÐ5vˆMØ(Ô:ˆLØ#0Ô#>Ð Ø -Ô 8Ðð  ,¨BÔ/°$Ô2JÑJˆL˜Ñð ð  ð  ð  å'*¨<¸Ô9TÑ'UÔ'Uð ñ  ô  Ðð
 Ð#¨Ð(<Ý œ?¨<¸¸¸¸1¸1¸1¸a¸a¸aÀ¸
Ô+CÍ5Ì9Ð]iÔ]pÐqÑqÔqˆLàÐ KÐ$7å œ;Ø˜A˜q !Ð+;¸BÔ+?Ô+EÐN^Ð_aÔNbÔNiðñ ô ˆLõ "œJ z°1°a½u¼{ÐScÐdfÔSgÔSnÐoÑoÔoÐoˆLàÐ"ð Ô    Ô%¨Ô)<Ô)LÒLÐLÝœmØ×%Ò%Ñ'Ô'ØÔ,Ô<Ø"'Ø#Ø"ðñ ô ÷ ’"�[Ô&Ñ'Ô'ð ð /3×.AÒ.AØ%Ø%Ø#Ø#ð	 /Bñ /
ô /
Ñ+ÐÐ+ð BSÀÔARð 
B
Ø-¨bÔ1Ø(CØ%6Ø$4Ø-Ø%5°c°r°cÔ%:Ø!5Ø-ð
B
ð 
B
ð ð
B
ð 
B
Ñ>Ð˜J¨Ð+>õ +Ø!Ø)Ø 3Ø-Ø!5Ø/ð
ñ 
ô 
ð 	
r.   c                 ó8  —  | j         |fddi|¤Ž}|j        }|j        }t          |¦  «        }| j                             |d         ¦  «        |d<   | j                             |d         ¦  «        |d<   d„ |D ¦   «         }d„ |D ¦   «         }||_        ||_        |S )zŠ
        pixel_values (`torch.FloatTensor`):
            Input pixel values of shape `(batch_size, num_channels, height, width)`.
        rB  Tr   r   c                 ób   — g | ],}|                      d ¦  «                             d dd¦  «        ‘Œ-S ©r   r   r   ©r|   rR   )rè   Úfeature_maps     r/   ré   z0Sam2Model.get_image_features.<locals>.<listcomp>L  s8   € Ð`Ð`Ð`ÀK˜×+Ò+¨AÑ.Ô.×6Ò6°q¸!¸QÑ?Ô?Ð`Ð`Ð`r.   c                 ób   — g | ],}|                      d ¦  «                             d dd¦  «        ‘Œ-S rZ  r[  )rè   Ú feature_maps_position_embeddingss     r/   ré   z0Sam2Model.get_image_features.<locals>.<listcomp>M  sH   € ð ,
ð ,
ð ,
à0ð -×4Ò4°QÑ7Ô7×?Ò?ÀÀ1ÀaÑHÔHð,
ð ,
ð ,
r.   )r5  r$   r%   r-  r8  r÷  rø  )rK   rS   rÁ   Úvision_outputsrK  r^  s         r/   rI  zSam2Model.get_image_features5  sÐ   € ð 3F°$Ô2EÀlÐ2oÐ2oÐ`dÐ2oÐhnÐ2oÐ2oˆà%Ô7ˆØ+9Ô+OÐ(õ ˜LÑ)Ô)ˆØÔ+×3Ò3°LÀ´OÑDÔDˆ�Q‰ØÔ+×3Ò3°LÀ´OÑDÔDˆ�Q‰ð aÐ`ÐS_Ð`Ñ`Ô`ˆð,
ð ,
à4Tð,
ñ ,
ô ,
Ð(ð ,8ˆÔ(Ø/OˆÔ,àÐr.   )NNNN)	NNNNNNTNN)!r&   r'   r(   r:  r   rÄ  ra  Ú_tied_weights_keysr   rB   rQ  r*   r   r@  r?  r+   r   r   r-  rL  Ú
LongTensorrP  r   r   r   r‹   r1   rX   r   r9   r#   rI  rY   rZ   s   @r/   r5  r5    så  ø€ € € € € ð )ÐØ4°n°nÐE]ÐefÐ6gÑ6gÔ6gÐhÐØÐð˜zð ð ð ð ð ð ð":ð :ð :ðB°e´lð Bð Bð Bð Bð €U„]�_„_ð àÔ'ð ð Ð+Ô,ð ð 
ˆeŒlÔ	ð	 ð  ð  ñ „_ð ð6 €U„]�_„_ð 26Ø04Ø04Ø/3ðð àÔ'¨$Ñ.ðð Ô&¨Ñ-ðð Ô&¨Ñ-ð	ð
 Ô%¨Ñ,ðð ð ñ „_ðð@  ØØð 26Ø15Ø04Ø04Ø/3Ø59Ø!%Ø9=Ø59ðk
ð k
àÔ'¨$Ñ.ðk
ð Ô'¨$Ñ.ðk
ð Ô&¨Ñ-ð	k
ð
 Ô&¨Ñ-ðk
ð Ô%¨Ñ,ðk
ð  Ô+¨dÑ2ðk
ð ðk
ð $Ô/°$Ñ6ðk
ð  Ô+¨dÑ2ðk
ð Ð+Ô,ðk
ð 
%ðk
ð k
ð k
ñ „^ñ „_ñ  Ôðk
ðZ ØðàÔ'ðð Ð+Ô,ðð 
Ð(Ñ	(ð	ð ð ñ „^ñ Ôðð ð ð ð r.   r5  )r5  rc  r#  r-  )r±   r‰   )Tre   Úcollections.abcr   Údataclassesr   ÚnumpyrG  r*   Útorch.nnrE   Útorch.nn.functionalr¾   r¨   r   Ú r   r/  Úactivationsr   Úmodeling_layersr	   Úmodeling_outputsr
   r   Úmodeling_utilsr   r   Úprocessing_utilsr   Úpytorch_utilsr   Úutilsr   r   r   r   Úutils.genericr   r   r   Úutils.output_capturingr   r   Úautor   Úconfiguration_sam2r   r   r   r   r    Ú
get_loggerr&   r¿  r#   r1   ÚModuler;   r\   r‘   rŒ   rÄ   rŠ   rÊ   rÌ   rÞ   rý   r  r  r  r#  r-  rc  r2  ry  rˆ  r³  rÄ  r×  r	  rƒ  ræ  r5  Ú__all__r-   r.   r/   ú<module>rv     so  ðð* €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø @Ð @Ð @Ð @Ð @Ð @Ø KÐ KÐ KÐ KÐ KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ iÐ iÐ iÐ iÐ iÐ iÐ iÐ iÐ iÐ iØ EÐ EÐ EÐ EÐ EÐ EÐ EÐ EØ Ð Ð Ð Ð Ð ðð ð ð ð ð ð ð ð ð ð ð ð ð ð 
ˆÔ	˜HÑ	%Ô	%€ð €ÐKÐLÑLÔLØ
ð;ð ;ð ;ð ;ð ;Ð8ñ ;ô ;ñ „ñ MÔLð;ð  €ÐFÐGÑGÔGØ
ðIð Ið Ið Ið I +ñ Iô Iñ „ñ HÔGðIð@ð ð ð ð ˜"œ)ñ ô ð ðBJ
ð J
ð J
ð J
ð J
 ¤	ñ J
ô J
ð J
ðZ18ð 18ð 18ð 18ð 18�R”Yñ 18ô 18ð 18ðv ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð ð%ð ð%ð %ð %ð %ð,ð ˆuŒ|ð ¨3°©:ð ÀÄð ð ð ð ð?ð ?ð ?ð ?ð ?˜bœiñ ?ô ?ð ?ðDð ð ð ð �b”iñ ô ð ð<2ð 2ð 2ð><ð <ð <ð:Yð Yð Yð Yð YÐ4ñ Yô Yð Yðx €ððñ ô ð
 ð<ð <ð <ð <ð <˜kñ <ô <ñ „ñô ð<ð ð8ð 8ð 8ð 8ð 8˜/ñ 8ô 8ñ „ð8ðBB
ð B
ð B
ð B
ð B
Ð+ñ B
ô B
ð B
ðJ €ððñ ô ð
/
ð /
ð /
ð /
ð /
Ð)ñ /
ô /
ñô ð
/
ðdSð Sð Sð Sð S˜bœiñ Sô Sð Sð2 ð  ð  ð  ð  ˜œ	ñ  ô  ð  ð6\3ð \3ð \3ð \3ð \3˜œ	ñ \3ô \3ð \3ð~D)ð D)ð D)ð D)ð D)�B”Iñ D)ô D)ð D)ðNJ'ð J'ð J'ð J'ð J'Ð9ñ J'ô J'ð J'ðZ6ð 6ð 6ð 6ð 6˜BœIñ 6ô 6ð 6ðrð ð ð ð �B”Lñ ô ð ð4H/ð H/ð H/ð H/ð H/�b”iñ H/ô H/ð H/ðV €ððñ ô ðrð rð rð rð rÐ#ñ rô rñô ðrðj	 WÐ
VÐ
V€€€r.   