Ë
    Fêñiåð  ã                  ó  — d dl mZ d dl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mZ ddlmZmZ dZ G d„ dej4                  «      Z G d„ dej                  j4                  «      Z G d„ de«      Zy)é    )ÚannotationsN)Únn)Útrunc_normal_)ÚMLP)ÚLOGGERé   )ÚSAM2TwoWayTransformerÚTwoWayTransformer)ÚMaskDecoderÚSAM2MaskDecoder)ÚImageEncoderViTÚPromptEncoder)Úget_1d_sine_peÚselect_closest_cond_framesg      �Àc                  óV   ‡ — e Zd ZU dZdZded<   	 	 d	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Zd„ Zˆ xZS )	ÚSAMModelaÆ  Segment Anything Model (SAM) for object segmentation tasks.

    This class combines image encoders, prompt encoders, and mask decoders to predict object masks from images and input
    prompts.

    Attributes:
        mask_threshold (float): Threshold value for mask prediction.
        image_encoder (ImageEncoderViT): Backbone for encoding images into embeddings.
        prompt_encoder (PromptEncoder): Encoder for various types of input prompts.
        mask_decoder (MaskDecoder): Predicts object masks from image and prompt embeddings.
        pixel_mean (torch.Tensor): Mean values for normalizing pixels in the input image.
        pixel_std (torch.Tensor): Standard deviation values for normalizing pixels in the input image.

    Methods:
        set_imgsz: Set image size to make model compatible with different image sizes.

    Examples:
        >>> image_encoder = ImageEncoderViT(...)
        >>> prompt_encoder = PromptEncoder(...)
        >>> mask_decoder = MaskDecoder(...)
        >>> sam_model = SAMModel(image_encoder, prompt_encoder, mask_decoder)
        >>> # Further usage depends on SAMPredictor class

    Notes:
        All forward() operations are implemented in the SAMPredictor class.
    ç        ÚfloatÚmask_thresholdc                ó(  •— t         ‰| �  «        || _        || _        || _        | j                  dt        j                  |«      j                  ddd«      d«       | j                  dt        j                  |«      j                  ddd«      d«       y)a¥  Initialize the SAMModel class to predict object masks from an image and input prompts.

        Args:
            image_encoder (ImageEncoderViT): The backbone used to encode the image into image embeddings.
            prompt_encoder (PromptEncoder): Encodes various types of input prompts.
            mask_decoder (MaskDecoder): Predicts masks from the image embeddings and encoded prompts.
            pixel_mean (list[float]): Mean values for normalizing pixels in the input image.
            pixel_std (list[float]): Standard deviation values for normalizing pixels in the input image.

        Notes:
            All forward() operations moved to SAMPredictor.
        Ú
pixel_meanéÿÿÿÿr   FÚ	pixel_stdN)	ÚsuperÚ__init__Úimage_encoderÚprompt_encoderÚmask_decoderÚregister_bufferÚtorchÚTensorÚview)Úselfr   r   r   r   r   Ú	__class__s         €úd/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/sam/modules/sam.pyr   zSAMModel.__init__7   s   ø€ ô( 	‰ÑÔØ*ˆÔØ,ˆÔØ(ˆÔØ×Ñ˜\¬5¯<©<¸
Ó+C×+HÑ+HÈÈQÐPQÓ+RÐTYÔZØ×Ñ˜[¬%¯,©,°yÓ*A×*FÑ*FÀrÈ1ÈaÓ*PÐRWÕXó    c                óþ   — t        | j                  d«      r| j                  j                  |«       || j                  _        |D �cg c]  }|dz  ‘Œ	 c}| j                  _        |d   | j                  _        yc c}w )úCSet image size to make model compatible with different image sizes.Ú	set_imgszé   r   N)Úhasattrr   r)   r   Úinput_image_sizeÚimage_embedding_sizeÚimg_size©r#   ÚimgszÚxs      r%   r)   zSAMModel.set_imgszR   si   € ä�4×%Ñ% {Ô3Ø×Ñ×(Ñ(¨Ô/Ø/4ˆ×ÑÔ,ØEJÖ3KÀ°A¸³GÒ3Kˆ×ÑÔ0Ø&+¨A¡hˆ×ÑÕ#ùò 4Ls   ÁA:))g33333ë^@gR¸…ë]@gR¸…ëáY@)gÃõ(\�2M@g�Âõ(\�L@g     °L@)r   r   r   r   r   r   r   úlist[float]r   r2   ÚreturnÚNone)	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ú__annotations__r   r)   Ú__classcell__©r$   s   @r%   r   r      sg   ø… ñð6  €N�EÓð #<Ø!8ðYà&ðYð &ðYð "ð	Yð
  ðYð ðYð 
õYö6/r&   r   c                  ó&  ‡ — e Zd ZU dZdZded<   	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Zed„ «       Zd„ Z	d„ Z
	 	 	 	 dd	„Zdd
„Zdd„Zdd„Z	 dd„Zd„ Zd„ Zd„ Z	 	 	 dd„Zd„ Zed„ «       Zdd„Zd„ Zˆ xZS )Ú	SAM2Modelaò  SAM2Model class for Segment Anything Model 2 with memory-based video object segmentation capabilities.

    This class extends the functionality of SAM to handle video sequences, incorporating memory mechanisms for temporal
    consistency and efficient tracking of objects across frames.

    Attributes:
        mask_threshold (float): Threshold value for mask prediction.
        image_encoder (ImageEncoderViT): Visual encoder for extracting image features.
        memory_attention (nn.Module): Module for attending to memory features.
        memory_encoder (nn.Module): Encoder for generating memory representations.
        num_maskmem (int): Number of accessible memory frames.
        image_size (int): Size of input images.
        backbone_stride (int): Stride of the backbone network output.
        sam_prompt_embed_dim (int): Dimension of SAM prompt embeddings.
        sam_image_embedding_size (int): Size of SAM image embeddings.
        sam_prompt_encoder (PromptEncoder): Encoder for processing input prompts.
        sam_mask_decoder (SAM2MaskDecoder): Decoder for generating object masks.
        obj_ptr_proj (nn.Module): Projection layer for object pointers.
        obj_ptr_tpos_proj (nn.Module): Projection for temporal positional encoding in object pointers.
        hidden_dim (int): Hidden dimension of the model.
        mem_dim (int): Memory dimension for encoding features.
        use_high_res_features_in_sam (bool): Whether to use high-resolution feature maps in the SAM mask decoder.
        use_obj_ptrs_in_encoder (bool): Whether to cross-attend to object pointers from other frames in the encoder.
        max_obj_ptrs_in_encoder (int): Maximum number of object pointers from other frames in encoder cross-attention.
        add_tpos_enc_to_obj_ptrs (bool): Whether to add temporal positional encoding to object pointers.
        proj_tpos_enc_in_obj_ptrs (bool): Whether to add an extra linear projection layer for temporal positional
            encoding in object pointers.
        use_signed_tpos_enc_to_obj_ptrs (bool): Whether to use signed distance in temporal positional encoding.
        only_obj_ptrs_in_the_past_for_eval (bool): Whether to only attend to object pointers in the past during
            evaluation.
        pred_obj_scores (bool): Whether to predict if there is an object in the frame.
        pred_obj_scores_mlp (bool): Whether to use an MLP to predict object scores.
        fixed_no_obj_ptr (bool): Whether to have a fixed no-object pointer when there is no object present.
        soft_no_obj_ptr (bool): Whether to mix in no-object pointer softly for easier recovery and error mitigation.
        use_mlp_for_obj_ptr_proj (bool): Whether to use MLP for object pointer projection.
        no_obj_embed_spatial (torch.Tensor | None): No-object embedding for spatial frames.
        max_cond_frames_in_attn (int): Maximum number of conditioning frames to participate in memory attention.
        directly_add_no_mem_embed (bool): Whether to directly add no-memory embedding to image feature on the first
            frame.
        multimask_output_in_sam (bool): Whether to output multiple masks for the first click on initial conditioning
            frames.
        multimask_min_pt_num (int): Minimum number of clicks to use multimask output in SAM.
        multimask_max_pt_num (int): Maximum number of clicks to use multimask output in SAM.
        multimask_output_for_tracking (bool): Whether to use multimask output for tracking.
        use_multimask_token_for_obj_ptr (bool): Whether to use multimask tokens for object pointers.
        iou_prediction_use_sigmoid (bool): Whether to use sigmoid to restrict IoU prediction to [0-1].
        memory_temporal_stride_for_eval (int): Memory bank's temporal stride during evaluation.
        non_overlap_masks_for_mem_enc (bool): Whether to apply non-overlapping constraints on object masks in memory
            encoder during evaluation.
        sigmoid_scale_for_mem_enc (float): Scale factor for mask sigmoid probability.
        sigmoid_bias_for_mem_enc (float): Bias factor for mask sigmoid probability.
        binarize_mask_from_pts_for_mem_enc (bool): Whether to binarize sigmoid mask logits on interacted frames with
            clicks during evaluation.
        use_mask_input_as_output_without_sam (bool): Whether to directly output the input mask without using SAM prompt
            encoder and mask decoder on frames with mask input.

    Methods:
        forward_image: Process image batch through encoder to extract multi-level features.
        track_step: Perform a single tracking step, updating object masks and memory features.
        set_binarize: Set binarize for VideoPredictor.
        set_imgsz: Set image size to make model compatible with different image sizes.

    Examples:
        >>> model = SAM2Model(image_encoder, memory_attention, memory_encoder)
        >>> image_batch = torch.rand(1, 3, 512, 512)
        >>> features = model.forward_image(image_batch)
        >>> track_results = model.track_step(0, True, features, None, None, None, {})
    r   r   r   c$                ód  •— t         ‰$| �  «        || _        || _        |rdnd| _        || _        || _        |r(t        j                  j                  dddd¬«      | _
        || _        |r|sJ ‚|| _        || _        || _        || _        |j                   | _        || _        | j"                  | _        t)        | j$                  d«      rRt)        | j$                  j*                  d«      r2| j$                  j*                  j,                  j.                  d   | _        || _        t        j                  j3                  t        j4                  |dd| j&                  «      «      | _        t9        | j6                  d¬	«       t        j                  j3                  t        j4                  dd| j"                  «      «      | _        t        j                  j3                  t        j4                  dd| j"                  «      «      | _        t9        | j:                  d¬	«       t9        | j<                  d¬	«       || _        || _         || _!        |	| _"        || _#        || _$        |
| _%        || _&        || _'        || _(        || _)        || _*        || _+        || _,        || _-        |"| _.        || _/        || _0        || _1        || _2        | jb                  r| j^                  sJ ‚| j
                  sJ ‚| j^                  re| j
                  rYt        j                  j3                  t        j4                  d| j"                  «      «      | _3        t9        | jf                  d¬	«       | | _4        d
| _5        |!rYt        j                  j3                  t        j4                  d| j&                  «      «      | _5        t9        | jj                  d¬	«       | jm                  «        || _7        d| _8        |#rRts        jt                  d«       t        jv                  | j                  jx                  ddd¬«      | j                  _<        y
y
)aÕ  Initialize the SAM2Model for video object segmentation with memory-based tracking.

        Args:
            image_encoder (nn.Module): Visual encoder for extracting image features.
            memory_attention (nn.Module): Module for attending to memory features.
            memory_encoder (nn.Module): Encoder for generating memory representations.
            num_maskmem (int): Number of accessible memory frames.
            image_size (int): Size of input images.
            backbone_stride (int): Stride of the image backbone output.
            sigmoid_scale_for_mem_enc (float): Scale factor for mask sigmoid probability.
            sigmoid_bias_for_mem_enc (float): Bias factor for mask sigmoid probability.
            binarize_mask_from_pts_for_mem_enc (bool): Whether to binarize sigmoid mask logits on interacted frames with
                clicks during evaluation.
            use_mask_input_as_output_without_sam (bool): Whether to directly output the input mask without using SAM
                prompt encoder and mask decoder on frames with mask input.
            max_cond_frames_in_attn (int): Maximum number of conditioning frames to participate in memory attention.
            directly_add_no_mem_embed (bool): Whether to directly add no-memory embedding to image feature on the first
                frame.
            use_high_res_features_in_sam (bool): Whether to use high-resolution feature maps in the SAM mask decoder.
            multimask_output_in_sam (bool): Whether to output multiple masks for the first click on initial conditioning
                frames.
            multimask_min_pt_num (int): Minimum number of clicks to use multimask output in SAM.
            multimask_max_pt_num (int): Maximum number of clicks to use multimask output in SAM.
            multimask_output_for_tracking (bool): Whether to use multimask output for tracking.
            use_multimask_token_for_obj_ptr (bool): Whether to use multimask tokens for object pointers.
            iou_prediction_use_sigmoid (bool): Whether to use sigmoid to restrict IoU prediction to [0-1].
            memory_temporal_stride_for_eval (int): Memory bank's temporal stride during evaluation.
            non_overlap_masks_for_mem_enc (bool): Whether to apply non-overlapping constraints on object masks in memory
                encoder during evaluation.
            use_obj_ptrs_in_encoder (bool): Whether to cross-attend to object pointers from other frames in the encoder.
            max_obj_ptrs_in_encoder (int): Maximum number of object pointers from other frames in encoder
                cross-attention.
            add_tpos_enc_to_obj_ptrs (bool): Whether to add temporal positional encoding to object pointers in the
                encoder.
            proj_tpos_enc_in_obj_ptrs (bool): Whether to add an extra linear projection layer for temporal positional
                encoding in object pointers.
            use_signed_tpos_enc_to_obj_ptrs (bool): Whether to use signed distance in the temporal positional encoding
                in the object pointers.
            only_obj_ptrs_in_the_past_for_eval (bool): Whether to only attend to object pointers in the past during
                evaluation.
            pred_obj_scores (bool): Whether to predict if there is an object in the frame.
            pred_obj_scores_mlp (bool): Whether to use an MLP to predict object scores.
            fixed_no_obj_ptr (bool): Whether to have a fixed no-object pointer when there is no object present.
            soft_no_obj_ptr (bool): Whether to mix in no-object pointer softly for easier recovery and error mitigation.
            use_mlp_for_obj_ptr_proj (bool): Whether to use MLP for object pointer projection.
            no_obj_embed_spatial (bool): Whether to add no-object embedding to spatial frames.
            sam_mask_decoder_extra_args (dict | None): Extra arguments for constructing the SAM mask decoder.
            compile_image_encoder (bool): Whether to compile the image encoder for faster inference.
        é   r   é   )Úkernel_sizeÚstrideÚout_projÚweightr   g{®Gáz”?)ÚstdNTzFImage encoder compilation is enabled. First forward pass will be slow.zmax-autotuneF)ÚmodeÚ	fullgraphÚdynamic)=r   r   r   Úuse_high_res_features_in_samÚnum_feature_levelsÚuse_obj_ptrs_in_encoderÚmax_obj_ptrs_in_encoderr    r   ÚConv2dÚmask_downsampleÚadd_tpos_enc_to_obj_ptrsÚproj_tpos_enc_in_obj_ptrsÚuse_signed_tpos_enc_to_obj_ptrsÚ"only_obj_ptrs_in_the_past_for_evalÚmemory_attentionÚd_modelÚ
hidden_dimÚmemory_encoderÚmem_dimr+   rC   rD   ÚshapeÚnum_maskmemÚ	ParameterÚzerosÚmaskmem_tpos_encr   Úno_mem_embedÚno_mem_pos_encÚdirectly_add_no_mem_embedÚsigmoid_scale_for_mem_encÚsigmoid_bias_for_mem_encÚ"binarize_mask_from_pts_for_mem_encÚnon_overlap_masks_for_mem_encÚmemory_temporal_stride_for_evalÚ$use_mask_input_as_output_without_samÚmultimask_output_in_samÚmultimask_min_pt_numÚmultimask_max_pt_numÚmultimask_output_for_trackingÚuse_multimask_token_for_obj_ptrÚiou_prediction_use_sigmoidÚ
image_sizeÚbackbone_strideÚsam_mask_decoder_extra_argsÚpred_obj_scoresÚpred_obj_scores_mlpÚfixed_no_obj_ptrÚsoft_no_obj_ptrÚ
no_obj_ptrÚuse_mlp_for_obj_ptr_projÚno_obj_embed_spatialÚ_build_sam_headsÚmax_cond_frames_in_attnÚ!add_all_frames_to_correct_as_condr   ÚinfoÚcompileÚforward©%r#   r   rS   rV   rY   rl   rm   r`   ra   rb   re   rw   r_   rI   rf   rg   rh   ri   rj   rk   rd   rc   rK   rL   rO   rP   rQ   rR   ro   rp   rq   rr   rt   ru   rn   Úcompile_image_encoderr$   s%                                       €r%   r   zSAM2Model.__init__£   su  ø€ ôn 	‰ÑÔð +ˆÔà,HˆÔ)Ù'C¡!ÈˆÔØ'>ˆÔ$Ø'>ˆÔ$Ù"ô $)§8¡8§?¡?°1°aÀQÈq ?Ó#QˆDÔ Ø(@ˆÔ%Ù$Ù+Ð+Ð+Ø)BˆÔ&Ø/NˆÔ,Ø2TˆÔ/ð !1ˆÔØ*×2Ñ2ˆŒð -ˆÔØ—‘ˆŒÜ�4×&Ñ&¨
Ô3¼À×@SÑ@S×@\Ñ@\Ð^fÔ8gà×.Ñ.×7Ñ7×>Ñ>×DÑDÀQÑGˆDŒLØ&ˆÔä %§¡× 2Ñ 2´5·;±;¸{ÈAÈqÐRV×R^ÑR^Ó3_Ó `ˆÔÜ�d×+Ñ+°Õ6ä!ŸH™H×.Ñ.¬u¯{©{¸1¸aÀÇÁÓ/QÓRˆÔÜ#Ÿh™h×0Ñ0´·±¸QÀÀ4Ç?Á?Ó1SÓTˆÔÜ�d×'Ñ'¨TÕ2Ü�d×)Ñ)¨tÕ4Ø)BˆÔ&ð *CˆÔ&Ø(@ˆÔ%Ø2TˆÔ/Ø-JˆÔ*Ø/NˆÔ,ð 5YˆÔ1Ø'>ˆÔ$Ø$8ˆÔ!Ø$8ˆÔ!Ø-JˆÔ*Ø/NˆÔ,Ø*DˆÔ'ð %ˆŒØ.ˆÔØ+FˆÔ(Ø.ˆÔØ#6ˆÔ Ø 0ˆÔØ.ˆÔØ× Ò Ø×'Ò'Ð'Ð'Ø×/Ò/Ð/Ð/Ø×Ò D×$@Ò$@Ü#Ÿh™h×0Ñ0´·±¸QÀÇÁÓ1PÓQˆDŒOÜ˜$Ÿ/™/¨tÕ4Ø(@ˆÔ%Ø$(ˆÔ!ÙÜ(-¯©×(:Ñ(:¼5¿;¹;ÀqÈ$Ï,É,Ó;WÓ(XˆDÔ%Ü˜$×3Ñ3¸Õ>à×ÑÔØ'>ˆÔ$Ø15ˆÔ.ñ !ä�K‰KÐ`ÔaÜ).¯©Ø×"Ñ"×*Ñ*Ø#ØØô	*ˆD×ÑÕ&ð !r&   c                óH   — t        | j                  «       «      j                  S )z=Return the device on which the model's parameters are stored.)ÚnextÚ
parametersÚdevice©r#   s    r%   r�   zSAM2Model.deviceY  s   € ô �D—O‘OÓ%Ó&×-Ñ-Ð-r&   c                ó   — t        d«      ‚)zWProcess image and prompt inputs to generate object masks and scores in video sequences.z„Please use the corresponding methods in SAM2VideoPredictor for inference.See notebooks/video_predictor_example.ipynb for an example.)ÚNotImplementedError)r#   ÚargsÚkwargss      r%   r{   zSAM2Model.forward^  s   € ä!ðJó
ð 	
r&   c                ó  — | j                   | _        | j                  | j                  z  | _        t        | j                  | j                  | j                  f| j                  | j                  fd¬«      | _        t        ddt        d| j                  dd¬«      | j                  dd| j                  | j                  | j                  | j                  | j                  d	œ
| j                  xs i ¤Ž| _        | j                   rwt"        j$                  j'                  | j                   | j                   «      | _        | j*                  rUt-        | j                   | j                   | j                   d«      | _        n#t"        j$                  j/                  «       | _        | j0                  r:t"        j$                  j'                  | j                   | j2                  «      | _        y
t"        j$                  j/                  «       | _        y
)zMBuild SAM-style prompt encoder and mask decoder for image segmentation tasks.r*   )Ú	embed_dimr-   r,   Úmask_in_chansr?   é   é   é   ©ÚdepthÚembedding_dimÚmlp_dimÚ	num_headsé   ©
Únum_multimask_outputsÚtransformerÚtransformer_dimÚiou_head_depthÚiou_head_hidden_dimÚuse_high_res_featuresrk   ro   rp   rj   N© )rU   Úsam_prompt_embed_dimrl   rm   Úsam_image_embedding_sizer   Úsam_prompt_encoderr   r	   rI   rk   ro   rp   rj   rn   Úsam_mask_decoderrK   r    r   ÚLinearÚobj_ptr_projrt   r   ÚIdentityrP   rW   Úobj_ptr_tpos_projr‚   s    r%   rv   zSAM2Model._build_sam_headse  s”  € à$(§O¡OˆÔ!Ø(,¯©¸4×;OÑ;OÑ(OˆÔ%ô #0Ø×/Ñ/à×-Ñ-Ø×-Ñ-ð"ð #Ÿo™o¨t¯©Ð?Øô#
ˆÔô !0ð !
Ø"#Ü-ØØ"×7Ñ7ØØô	ð !×5Ñ5ØØ #Ø"&×"CÑ"CØ'+×'FÑ'FØ ×0Ñ0Ø $× 8Ñ 8Ø,0×,PÑ,Pñ!
ð  ×/Ñ/Ò5°2ñ!!
ˆÔð$ ×'Ò'ä %§¡§¡°·±ÀÇÁÓ QˆDÔØ×,Ò,Ü$'¨¯©¸¿¹È$Ï/É/Ð[\Ó$]�Õ!ä %§¡× 1Ñ 1Ó 3ˆDÔØ×)Ò)ô &+§X¡X§_¡_°T·_±_ÀdÇlÁlÓ%SˆDÕ"ä%*§X¡X×%6Ñ%6Ó%8ˆDÕ"r&   c           	     ó®  — |j                   d   }|j                  }|j                  d«      | j                  k(  sJ ‚|j                  d«      | j                  k(  sJ ‚|j                  d«      | j                  k(  sJ ‚|�0|d   }|d   }	|j                   d   |k(  r|	j                   d   |k(  sNJ ‚t        j                  |dd||j                  ¬«      }t        j                  |dt
        j                  |¬	«       }	|�Ÿt        |j                   «      d
k(  r|j                   dd |dfk(  sJ ‚|j                   dd | j                  j                  k7  rHt        j                  |j                  |j                  «      | j                  j                  ddd¬«      }
n|}
nd}
| j                  ||	fd|
¬«      \  }}| j!                  || j                  j#                  «       |||d|¬«      \  }}}}| j$                  r(|dkD  }t        j&                  |dd…ddf   |t(        «      }t        j                  || j*                  | j*                  fdd¬«      }|dd…df   }|rvt        j,                  |d¬«      }t        j.                  ||¬«      }|||f   j1                  d«      }|||f   j1                  d«      }|j                  d«      dkD  r|||f   }n||}}| j3                  |«      }| j$                  r^| j4                  r|j7                  «       }nj                  |j                  «      }| j8                  r||z  }|d|z
  | j:                  z  z   }|||||||fS )a[
  Forward pass through SAM prompt encoders and mask heads.

        This method processes image features and optional point/mask inputs to generate object masks and scores.

        Args:
            backbone_features (torch.Tensor): Image features with shape (B, C, H, W).
            point_inputs (dict[str, torch.Tensor] | None): Dictionary containing point prompts with keys 'point_coords'
                (Tensor of shape (B, P, 2) with float32 dtype, containing absolute pixel-unit coordinates in (x, y)
                format for P input points) and 'point_labels' (Tensor of shape (B, P) with int32 dtype, where 1 means
                positive clicks, 0 means negative clicks, and -1 means padding).
            mask_inputs (torch.Tensor | None): Mask of shape (B, 1, H*16, W*16), float or bool, with the same spatial
                size as the image.
            high_res_features (list[torch.Tensor] | None): List of two feature maps with shapes (B, C, 4*H, 4*W) and (B,
                C, 2*H, 2*W) respectively, used as high-resolution feature maps for SAM decoder.
            multimask_output (bool): If True, output 3 candidate masks and their IoU estimates; if False, output only 1
                mask and its IoU estimate.

        Returns:
            low_res_multimasks (torch.Tensor): Tensor of shape (B, M, H*4, W*4) with SAM output mask logits.
            high_res_multimasks (torch.Tensor): Tensor of shape (B, M, H*16, W*16) with upsampled mask logits.
            ious (torch.Tensor): Tensor of shape (B, M) with estimated IoU for each output mask.
            low_res_masks (torch.Tensor): Tensor of shape (B, 1, H*4, W*4) with the best low-resolution mask.
            high_res_masks (torch.Tensor): Tensor of shape (B, 1, H*16, W*16) with the best high-resolution mask.
            obj_ptr (torch.Tensor): Tensor of shape (B, C) with object pointer vector for the output mask.
            object_score_logits (torch.Tensor): Tensor of shape (B, 1) with object score logits.

        Examples:
            >>> backbone_features = torch.rand(1, 256, 32, 32)
            >>> point_inputs = {"point_coords": torch.rand(1, 2, 2), "point_labels": torch.tensor([[1, 0]])}
            >>> mask_inputs = torch.rand(1, 1, 512, 512)
            >>> results = model._forward_sam_heads(backbone_features, point_inputs, mask_inputs)
            >>> (
            ...     low_res_multimasks,
            ...     high_res_multimasks,
            ...     ious,
            ...     low_res_masks,
            ...     high_res_masks,
            ...     obj_ptr,
            ...     object_score_logits,
            ... ) = results
        r   r   rŠ   r?   NÚpoint_coordsÚpoint_labels©r�   Údtype)r§   r�   r@   éþÿÿÿFÚbilinearT©ÚsizeÚalign_cornersrF   Ú	antialias)ÚpointsÚboxesÚmasks)Úimage_embeddingsÚimage_peÚsparse_prompt_embeddingsÚdense_prompt_embeddingsÚmultimask_outputÚrepeat_imageÚhigh_res_features)r«   rF   r¬   r   ©Údim©r�   )rX   r�   r«   r›   rœ   r    r[   r§   ÚonesÚint32Úlenr�   Úmask_input_sizeÚFÚinterpolateÚtorž   Úget_dense_pero   ÚwhereÚNO_OBJ_SCORErl   ÚargmaxÚarangeÚ	unsqueezer    rr   Úsigmoidrq   rs   )r#   Úbackbone_featuresÚpoint_inputsÚmask_inputsr·   rµ   ÚBr�   Úsam_point_coordsÚsam_point_labelsÚsam_mask_promptÚsparse_embeddingsÚdense_embeddingsÚlow_res_multimasksÚiousÚsam_output_tokensÚobject_score_logitsÚis_obj_appearingÚhigh_res_multimasksÚsam_output_tokenÚbest_iou_indsÚ
batch_indsÚlow_res_masksÚhigh_res_masksÚobj_ptrÚlambda_is_obj_appearings                             r%   Ú_forward_sam_headszSAM2Model._forward_sam_heads”  s®  € ðb ×#Ñ# AÑ&ˆØ"×)Ñ)ˆØ ×%Ñ% aÓ(¨D×,EÑ,EÒEÐEÐEØ ×%Ñ% aÓ(¨D×,IÑ,IÒIÐIÐIØ ×%Ñ% aÓ(¨D×,IÑ,IÒIÐIÐIð Ð#Ø+¨NÑ;ÐØ+¨NÑ;ÐØ#×)Ñ)¨!Ñ,°Ò1Ð6F×6LÑ6LÈQÑ6OÐSTÒ6TÐTÐTô  %Ÿ{™{¨1¨a°¸6ÐIZ×I`ÑI`ÔaÐÜ %§
¡
¨1¨a´u·{±{È6Ô RÐRÐð Ð"ô �{×(Ñ(Ó)¨QÒ.°;×3DÑ3DÀRÀaÐ3HÈQÐPQÈFÒ3RÐRÐRØ× Ñ   Ð%¨×)@Ñ)@×)PÑ)PÒPÜ"#§-¡-Ø—N‘NÐ#4×#:Ñ#:Ó;Ø×0Ñ0×@Ñ@Ø"'Ø#Ø"ô#‘ð #.‘ð #ˆOà.2×.EÑ.EØ$Ð&6Ð7ØØ!ð /Fó /
Ñ+ÐÐ+ð
 LP×K`ÑK`Ø.Ø×,Ñ,×9Ñ9Ó;Ø%6Ø$4Ø-ØØ/ð Laó L
ÑHÐ˜DÐ"3Ð5Hð ×ÒØ2°QÑ6Ðô "'§¡Ð-=ºaÀÀt¸mÑ-LÐN`ÔbnÓ!oÐô  Ÿm™mØØ—/‘/ 4§?¡?Ð3ØØô	
Ðð -ªQ°¨TÑ2ÐÙä!ŸL™L¨°2Ô6ˆMÜŸ™ a°Ô7ˆJØ.¨z¸=Ð/HÑI×SÑSÐTUÓVˆMØ0°¸]Ð1JÑK×UÑUÐVWÓXˆNØ ×%Ñ% aÓ(¨1Ò,Ø#4°ZÀÐ5NÑ#OÑ à,>Ð@S˜>ˆMð ×#Ñ#Ð$4Ó5ˆØ×Òà×#Ò#Ø*=×*EÑ*EÓ*GÑ'à*:×*=Ñ*=¸g¿m¹mÓ*LÐ'à×$Ò$Ø1°GÑ;�Ø Ð%<Ñ!<ÀÇÁÑ OÑOˆGàØØØØØØð
ð 	
r&   c                óP  — d\  }}|j                  «       }||z  |z   }t        j                  ||j                  d«      dz  |j                  d«      dz  fddd¬«      }|j	                  |j
                  d	   d
«      j                  «       }	| j                  r|�|€:t        j                  |j
                  d	   | j                  |j                  ¬«      }
nD| j                  || j                  |j                  |j                  «      «      |¬«      \  }}}}}}
}t        j                  |j!                  d
«      j                  «       dkD  d
¬«      }|d   }|j                  «       }||z  |z   }| j"                  r&| j$                  r||
z  }
|
d
|z
  | j&                  z  z   }
|||	|||
|fS )zFProcess mask inputs directly as output, bypassing SAM encoder/decoder.)g      4@ç      $Àr¨   r@   r   Fr©   Trª   r   r   rº   )rÉ   rË   r·   r   r¸   ).N)r   r¿   rÀ   r«   Únew_onesrX   rK   r    r[   rU   r�   rß   rN   rÁ   r§   ÚanyÚflattenro   rq   rs   )r#   rË   rÉ   r·   Ú	out_scaleÚout_biasÚmask_inputs_floatrÜ   rÛ   rÓ   rÝ   Ú_rÖ   rÞ   rÕ   s                  r%   Ú_use_mask_as_outputzSAM2Model._use_mask_as_output(  sÎ  € ð *Ñˆ	�8Ø'×-Ñ-Ó/ÐØ*¨YÑ6¸ÑAˆÜŸ™ØØ ×%Ñ% bÓ)¨QÑ.°×0CÑ0CÀBÓ0GÈ1Ñ0LÐMØØØô
ˆð ×#Ñ# K×$5Ñ$5°aÑ$8¸!Ó<×BÑBÓDˆØ×+Ò+Ð/@Ð/HÐL]ÐLeä—k‘k +×"3Ñ"3°AÑ"6¸¿¹ÐP[×PbÑPbÔc‰Gð )-×(?Ñ(?Ø"3Ø ×0Ñ0Ð1B×1EÑ1EÐFW×F]ÑF]Ó1^Ó_Ø"3ð )@ó )Ñ%ˆAˆq�!�Q˜˜7 Aô !Ÿ9™9 [×%8Ñ%8¸Ó%;×%AÑ%AÓ%CÀcÑ%IÈqÔQÐØ+¨IÑ6ÐØ"2×"8Ñ"8Ó":ÐØ'Ð*AÑAÀHÑLÐØ×ÒØ×$Ò$Ø1°GÑ;�Ø Ð%<Ñ!<ÀÇÁÑ OÑOˆGð ØØØØØØð
ð 	
r&   c                óÜ   — | j                  |«      }| j                  rN| j                  j                  |d   d   «      |d   d<   | j                  j	                  |d   d   «      |d   d<   |S ©zRProcess image batch through encoder to extract multi-level features for SAM model.Úbackbone_fpnr   r   )r   rI   rž   Úconv_s0Úconv_s1©r#   Ú	img_batchÚbackbone_outs      r%   Úforward_imagezSAM2Model.forward_imageW  s{   € à×)Ñ)¨)Ó4ˆØ×,Ò,ð /3×.CÑ.C×.KÑ.KÈLÐYgÑLhÐijÑLkÓ.lˆL˜Ñ(¨Ñ+Ø.2×.CÑ.C×.KÑ.KÈLÐYgÑLhÐijÑLkÓ.lˆL˜Ñ(¨Ñ+ØÐr&   c                ó¾  — |dkD  rOi |¥|d   D �cg c]  }|j                  |ddd«      ‘Œ c}|d   D �cg c]  }|j                  |ddd«      ‘Œ c}dœ¥}t        |d   «      t        |d   «      k(  sJ ‚t        |d   «      | j                  k\  sJ ‚|d   | j                   d }|d   | j                   d }|D �cg c]   }|j                  d   |j                  d   f‘Œ" }}|D �cg c]$  }|j	                  d«      j                  dd	d«      ‘Œ& }	}|D �cg c]$  }|j	                  d«      j                  dd	d«      ‘Œ& }}||	||fS c c}w c c}w c c}w c c}w c c}w )
zZPrepare and flatten visual features from the image backbone output for further processing.r   rì   r   Úvision_pos_enc)rì   rô   Nr¨   rŠ   r   )Úexpandr½   rJ   rX   rä   Úpermute)
r#   rñ   ÚbatchÚfeatÚposÚfeature_mapsÚvision_pos_embedsr1   Ú
feat_sizesÚvision_featss
             r%   Ú_prepare_backbone_featuresz$SAM2Model._prepare_backbone_featuresa  s€  € à�1Š9ðØðàLXÐYgÑLhÖ iÀD §¡¨U°B¸¸BÕ!?Ò iØLXÐYiÑLjÖ"kÀS 3§:¡:¨e°R¸¸RÕ#@Ò"kòˆLô
 �< Ñ/Ó0´C¸ÐEUÑ8VÓ4WÒWÐWÐWÜ�< Ñ/Ó0°D×4KÑ4KÒKÐKÐKà# NÑ3°T×5LÑ5LÐ4LÐ4NÐOˆØ(Ð)9Ñ:¸D×<SÑ<SÐ;SÐ;UÐVÐà:KÖL°Q�q—w‘w˜r‘{ A§G¡G¨B¡KÒ0ÐLˆ
ÐLà?KÖL¸!˜Ÿ	™	 !›×,Ñ,¨Q°°1Õ5ÐLˆÐLØDUÖV¸q˜QŸY™Y q›\×1Ñ1°!°Q¸Õ:ÐVÐÐVØ˜\Ð+<¸jÐHÐHùò !jùÚ"kùò MùâLùÚVs   �E´EÂ;%EÃ&)EÄ)Ec	                óV  — |d   j                  d«      }	| j                  }
|d   \  }}|d   j                  }| j                  dk(  r(|d   j	                  ddd«      j                  |	|
||«      S d}|rdnd}|�s g g }}t        |d   «      dkD  sJ ‚|d   }t        ||| j                  «      \  }}|j                  «       D �cg c]  }d|f‘Œ }}| j                  rdn| j                  }t        d| j                  «      D ]�  }| j                  |z
  }|dk(  r|r||z   n||z
  }n1|s|dz
  |z  |z  }||dz
  |z  z
  }n|dz    |z   |z  }||dz
  |z  z   }|d   j                  |d«      }|€|j                  |d«      }|j                  ||f«       Œ’ |D ]É  \  }}|€Œ	|d   j                  ||j                   d	k(  ¬
«      }|j                  |j#                  d«      j	                  ddd«      «       |d   d   j                  |¬«      }|j#                  d«      j	                  ddd«      }|| j$                  | j                  |z
  dz
     z   }|j                  |«       ŒË | j&                  �rBt)        || j*                  «      }| j                  s=| j,                  r1|j/                  «       D ��ci c]  \  }}|r||k\  r	n||k  r||“Œ } }}n|} | j/                  «       D ��cg c],  \  }}| j0                  r||z
  |z  nt3        ||z
  «      |d   f‘Œ. }!}}t        d|«      D ]Z  }"|r||"z   n||"z
  }|dk  s|�||k\  r n@|d   j                  ||j                  |d«      «      }|€ŒE|!j                  |"|d   f«       Œ\ |!�r–t5        |!Ž \  }#}$t7        j8                  |$d¬«      }%| j:                  r’|dz
  }&| j<                  r|
n| j>                  }'t7        j@                  |#||d   jB                  ¬«      }(tE        |(|&z  |'¬«      }(| jG                  |(«      }(|(jI                  d«      jK                  d|	| j>                  «      }(n&|%jM                  t        |#«      |	| j>                  «      }(| j>                  |
k  ro|%jO                  d|	|
| j>                  z  | j>                  «      }%|%j	                  dddd«      j#                  dd«      }%|(jQ                  |
| j>                  z  d¬«      }(|j                  |%«       |j                  |(«       |%jR                  d   }n˜d}n•| jT                  r9|d   | jV                  z   })|)j	                  ddd«      j                  |	|
||«      })|)S | jV                  jK                  d|	| j>                  «      g}| jX                  jK                  d|	| j>                  «      g}t7        jZ                  |d¬«      }*t7        jZ                  |d¬«      }+| j]                  |||*|+|¬«      })|)j	                  ddd«      j                  |	|
||«      })|)S c c}w c c}}w c c}}w )zePrepare memory-conditioned features by fusing current frame's visual features with previous memories.r   r   r   rŠ   Úcond_frame_outputsÚnon_cond_frame_outputsNÚmaskmem_featuresÚcuda)r�   Únon_blockingÚmaskmem_pos_encrº   rÝ   r¸   r¦   r?   )ÚcurrÚcurr_posÚmemoryÚ
memory_posÚnum_obj_ptr_tokens)/r«   rU   r�   rY   rö   r"   r½   r   rw   ÚvaluesÚtrainingrd   ÚrangeÚgetÚappendrÁ   Útyperä   r\   rK   ÚminrL   rR   ÚitemsrQ   ÚabsÚzipr    ÚstackrO   rP   rW   Útensorr§   r   r¢   rÇ   rõ   Ú	new_zerosÚreshapeÚrepeat_interleaverX   r_   r]   r^   ÚcatrS   ),r#   Ú	frame_idxÚis_init_cond_frameÚcurrent_vision_featsÚcurrent_vision_pos_embedsrü   Úoutput_dictÚ
num_framesÚtrack_in_reverserÌ   ÚCÚHÚWr�   r
  Útpos_sign_mulÚto_cat_memoryÚto_cat_memory_pos_embedÚcond_outputsÚselected_cond_outputsÚunselected_cond_outputsÚoutÚt_pos_and_prevsÚrÚt_posÚt_relÚprev_frame_idxÚprevÚfeatsÚmaskmem_encrL   ÚtÚptr_cond_outputsÚpos_and_ptrsÚt_diffÚpos_listÚ	ptrs_listÚobj_ptrsÚ
t_diff_maxÚtpos_dimÚobj_posÚpix_feat_with_memr  Úmemory_pos_embeds,                                               r%   Ú$_prepare_memory_conditioned_featuresz.SAM2Model._prepare_memory_conditioned_featuresu  s·  € ð ! Ñ$×)Ñ)¨!Ó,ˆØ�O‰OˆØ˜"‰~‰ˆˆ1Ø% bÑ)×0Ñ0ˆð ×Ñ˜qÒ Ø'¨Ñ+×3Ñ3°A°q¸!Ó<×AÑAÀ!ÀQÈÈ1ÓMÐMØÐÙ.™°Aˆâ!à57¸Ð2ˆMô �{Ð#7Ñ8Ó9¸AÒ=Ð=Ð=à&Ð';Ñ<ˆLÜ=WØ˜<¨×)EÑ)Eó>Ñ:Ð!Ð#:ð 4I×3OÑ3OÓ3QÖR¨C  3šxÐRˆOÐRð
 —]’]‘¨×(LÑ(LˆAÜ˜q $×"2Ñ"2Ó3ò 5�Ø×(Ñ(¨5Ñ0�Ø˜A’:á:J Y°Ò%6ÐPYÐ\aÑPa‘NÙ)ð (1°1¡}¸Ñ&:¸aÑ%?�Nà%3°u¸q±yÀA±oÑ%E‘Nð *3°Q©Ð'7¸1Ñ'<Ð%=ÀÑ%A�Nà%3°u¸q±yÀA±oÑ%E�NØ!Ð":Ñ;×?Ñ?ÀÐPTÓU�Ø�;ð 2×5Ñ5°nÀdÓK�CØ×&Ñ&¨¨s |Õ4ð-5ð0  /ò <‘��tØ�<Øð Ð/Ñ0×3Ñ3¸6ÐPV×P[ÑP[Ð_eÑPeÐ3Óf�Ø×$Ñ$ U§]¡]°1Ó%5×%=Ñ%=¸aÀÀAÓ%FÔGà"Ð#4Ñ5°bÑ9×<Ñ<ÀFÐ<ÓK�Ø)×1Ñ1°!Ó4×<Ñ<¸QÀÀ1ÓE�à)¨D×,AÑ,AÀ$×BRÑBRÐUZÑBZÐ]^ÑB^Ñ,_Ñ_�Ø'×.Ñ.¨{Õ;ð<ð ×+Ó+Ü*-¨j¸$×:VÑ:VÓ*WÐ'ð —}’}¨×)PÒ)Pð '<×&AÑ&AÓ&C÷(á"˜A˜sÙ.>˜A ›NÀAÈÂNð ˜3™ð(Ð$ò (ð (=Ð$ð #3×"8Ñ"8Ó":÷ ñ ˜˜3ð  $×CÒCð '¨™]¨mÒ;ä!$ Y°¡]Ó!3à˜I™òð �ñ  ô $ AÐ'>Ó?ò F�FÙ.>˜	 FÒ*ÀIÐPVÑDV�AØ˜1’u Ð!7¸AÀºOÙØ%Ð&>Ñ?×CÑCÀAÐG^×GbÑGbÐcdÐfjÓGkÓl�CØ‘Ø$×+Ñ+¨V°S¸±^Ð,DÕEðFò  Ü*-¨|Ð*<Ñ'�H˜iä$Ÿ{™{¨9¸!Ô<�Hð ×4Ò4Ø%<¸qÑ%@˜
Ø(,×(FÒ(F¡1ÈDÏLÉL˜Ü"'§,¡,¨xÀÐNbÐceÑNf×NlÑNlÔ"m˜Ü"0°¸:Ñ1EÈ8Ô"T˜Ø"&×"8Ñ"8¸Ó"A˜Ø")×"3Ñ"3°AÓ"6×"=Ñ"=¸bÀ!ÀTÇ\Á\Ó"R™à"*×"4Ñ"4´S¸³]ÀAÀtÇ|Á|Ó"T˜Ø—|‘| aÒ'à#+×#3Ñ#3°B¸¸1ÀÇÁÑ;LÈdÏlÉlÓ#[˜Ø#+×#3Ñ#3°A°q¸!¸QÓ#?×#GÑ#GÈÈ1Ó#M˜Ø")×";Ñ";¸AÀÇÁÑ<MÐSTÐ";Ó"U˜Ø!×(Ñ(¨Ô2Ø+×2Ñ2°7Ô;Ø)1¯©¸Ñ):Ñ&à)*Ñ&ð ×-Ò-à$8¸Ñ$<¸t×?PÑ?PÑ$PÐ!Ø$5×$=Ñ$=¸aÀÀAÓ$F×$KÑ$KÈAÈqÐRSÐUVÓ$WÐ!Ø(Ð(ð "×.Ñ.×5Ñ5°a¸¸D¿L¹LÓIÐJˆMØ'+×':Ñ':×'AÑ'AÀ!ÀQÈÏÉÓ'UÐ&VÐ#ô —‘˜=¨aÔ0ˆÜ Ÿ9™9Ð%<À!ÔDÐà ×1Ñ1Ø%Ø.ØØ'Ø1ð 2ó 
Ðð .×5Ñ5°a¸¸AÓ>×CÑCÀAÀqÈ!ÈQÓOÐØ Ð ùòA Sùód(ùó s   ÃXÊ.XË!1X%c                óò  — |d   j                  d«      }| j                  }|d   \  }}	|d   j                  ddd«      j                  ||||	«      }
| j                  r| j
                  s| j                  |«      }| j                  xr |}|r+| j
                  s|dkD  j                  |
j                  «      }nt        j                  |«      }| j                  dk7  r|| j                  z  }| j                  dk7  r|| j                  z   }| j                  |
|d¬«      }|d	   }| j                  �E|dkD  j!                  «       }|d|d
   z
   | j                  d
   j"                  |j$                  Ž z  z  }||d   fS )zXEncode frame features and masks into a new memory representation for video segmentation.r   r   rŠ   r   ç      ð?r   T)Úskip_mask_sigmoidÚvision_features©.NNrô   )r«   rU   rö   r"   rc   r  Ú"_apply_non_overlapping_constraintsrb   rÁ   r§   r    rÈ   r`   ra   rV   ru   r   rõ   rX   )r#   r  rü   Úpred_masks_high_resrÕ   Úis_mask_from_ptsrÌ   r"  r#  r$  Úpix_featÚbinarizeÚmask_for_memÚmaskmem_outr  rÖ   s                   r%   Ú_encode_new_memoryzSAM2Model._encode_new_memory  s�  € ð ! Ñ$×)Ñ)¨!Ó,ˆØ�O‰OˆØ˜"‰~‰ˆˆ1à'¨Ñ+×3Ñ3°A°q¸!Ó<×AÑAÀ!ÀQÈÈ1ÓMˆØ×-Ò-°d·m²mð #'×"IÑ"IÐJ]Ó"^Ðà×:Ñ:ÒOÐ?OˆÙ˜DŸMšMØ/°!Ñ3×7Ñ7¸¿¹ÓG‰Lô !Ÿ=™=Ð)<Ó=ˆLà×)Ñ)¨SÒ0Ø'¨$×*HÑ*HÑHˆLØ×(Ñ(¨CÒ/Ø'¨$×*GÑ*GÑGˆLØ×)Ñ)¨(°LÐTXÐ)ÓYˆØ&Ð'8Ñ9Ðð ×$Ñ$Ð0Ø 3°aÑ 7×>Ñ>Ó@ÐØ Ð%5°oÑ%FÑ!Fð KÈ$×JcÑJcØñKç‰fÐ&×,Ñ,ðK.ñ !.ñ .Ðð   Ð-=Ñ!>Ð>Ð>r&   c           
     ó^  — t        |«      dkD  rft        |dd |dd «      D ��cg c]H  \  }} |j                  ddd«      j                  |j	                  d«      |j	                  d«      g|¢­Ž ‘ŒJ }}}nd}|�W| j
                  rK|d   j                  ddd«      } |j                  d| j                  g|d   ¢­Ž }| j                  |||«      }nT| j                  |||dd |dd |dd ||	|
¬«      }|�|�|�J ‚|}| j                  ||«      }| j                  |||||¬«      }|||fS c c}}w )úhPerform a single tracking step, updating object masks and memory features based on current frame inputs.r   Nr   rŠ   r   )r  r  r  r  rü   r  r   r!  )rÉ   rÊ   rË   r·   rµ   )r½   r  rö   r"   r«   re   rU   ré   r@  Ú_use_multimaskrß   )r#   r  r  r  r  rü   rÊ   rË   r  r   r!  Úprev_sam_mask_logitsr1   Úsr·   rI  Úsam_outputsrµ   s                     r%   Ú_track_stepzSAM2Model._track_stepD  s‘  € ô  Ð#Ó$ qÒ(ô  Ð 4°S°bÐ 9¸:ÀcÀr¸?ÓK÷!á�A�qð (�—	‘	˜!˜Q Ó"×'Ñ'¨¯©¨q«	°1·6±6¸!³9ÐA¸qÔAð!Ðò !ð
 !%ÐØÐ" t×'PÒ'Pð ,¨BÑ/×7Ñ7¸¸1¸aÓ@ˆHØ$�x—}‘} R¨¯©ÐJ¸:Àb¹>ÒJˆHØ×2Ñ2°;ÀÐJ[Ó\‰Kð ×@Ñ@Ø#Ø#5Ø%9¸"¸#Ð%>Ø*CÀBÀCÐ*HØ% b c˜?Ø'Ø%Ø!1ð Aó 	ˆHð $Ð/Ø#Ð/°KÐ4GÐGÐGØ2�Ø#×2Ñ2Ð3EÀ|ÓTÐØ×1Ñ1Ø"*Ø)Ø'Ø"3Ø!1ð 2ó ˆKð Ð-¨xÐ7Ð7ùóO!s   ¤AD)c                ó†   — |r5| j                   dkD  r&| j                  |||||du¬«      \  }}	||d<   |	|d<   yd|d<   d|d<   y)z^Run memory encoder on predicted mask to encode it into a new memory feature for future frames.r   N)r  rü   rG  rÕ   rH  r  r  )rY   rM  )
r#   r  rü   rÊ   Úrun_mem_encoderrÜ   rÕ   Úcurrent_outr  r  s
             r%   Ú_encode_memory_in_outputz"SAM2Model._encode_memory_in_output~  sr   € ñ ˜t×/Ñ/°!Ò3Ø04×0GÑ0GØ%9Ø%Ø$2Ø$7Ø".°dÐ":ð 1Hó 1Ñ-Ð˜oð /?ˆKÐ*Ñ+Ø-<ˆKÐ)Ò*à.2ˆKÐ*Ñ+Ø-1ˆKÐ)Ò*r&   c                ó´   — | j                  |||||||||	|
|«      \  }}}|\  }}}}}}}|||dœ}| j                  s||d<   | j                  |||||||«       |S )rO  )Ú
pred_masksrG  rÝ   rÕ   )rT  r  rX  )r#   r  r  r  r  rü   rÊ   rË   r  r   r!  rV  rQ  rS  rè   rÛ   rÜ   rÝ   rÕ   rW  s                       r%   Ú
track_stepzSAM2Model.track_step—  s«   € ð, !×,Ñ,ØØØ Ø%ØØØØØØØ ó
Ñˆ�Q˜ð P[ÑLˆˆ1ˆa� °Ð9Lð (Ø#1Øñ
ˆð
 �}Š}ð 2EˆKÐ-Ñ.ð 	×%Ñ%Ø ØØØØØØô	
ð Ðr&   c                óº   — |€dn|d   j                  d«      }| j                  xr6 |xs | j                  xr$ | j                  |cxk  xr | j                  k  S c S )zaDetermine whether to use multiple mask outputs in the SAM head based on configuration and inputs.r   r¥   r   )r«   rf   ri   rg   rh   )r#   r  rÊ   Únum_ptss       r%   rP  zSAM2Model._use_multimaskÓ  sk   € à#Ð+‘!°¸nÑ1M×1RÑ1RÐSTÓ1Uˆà×(Ñ(ò TØ#ÒI t×'IÑ'IòTà×*Ñ*¨gÖR¸×9RÑ9RÑRð	
ñ Sð	
r&   c                ó  — | j                   d   }|dk(  r| S | j                  }t        j                  | dd¬«      }t        j                  ||¬«      dd…dddf   }||k(  }t        j
                  || t        j                  | d¬«      «      } | S )	z\Apply non-overlapping constraints to masks, keeping the highest scoring object per location.r   r   T)r¹   Úkeepdimrº   Nrá   ©Úmax)rX   r�   r    rÅ   rÆ   rÃ   Úclamp)rZ  Ú
batch_sizer�   Úmax_obj_indsÚbatch_obj_indsÚkeeps         r%   rF  z,SAM2Model._apply_non_overlapping_constraintsÜ  sŒ   € ð  ×%Ñ% aÑ(ˆ
Ø˜Š?ØÐà×"Ñ"ˆä—|‘| J°A¸tÔDˆäŸ™ j¸Ô@ÂÀDÈ$ÐPTÐATÑUˆØ˜~Ñ-ˆô —[‘[  z´5·;±;¸zÈuÔ3UÓVˆ
ØÐr&   c                ó   — || _         y)z Set binarize for VideoPredictor.N)rb   )r#   rJ  s     r%   Úset_binarizezSAM2Model.set_binarizeî  s
   € à2:ˆÕ/r&   c                ó¢  — t        | j                  d«      r| j                  j                  |«       |d   | _        || j                  _        |D �cg c]  }|| j                  z  ‘Œ c}| j                  _        |D �cg c]  }|| j                  z  dz  ‘Œ c}| j                  _        | j                  | j                  z  | _	        yc c}w c c}w )r(   r)   r   r@   N)
r+   r   r)   rl   r�   r,   rm   r-   r¾   rœ   r/   s      r%   r)   zSAM2Model.set_imgszò  s»   € ä�4×%Ñ% {Ô3Ø×Ñ×(Ñ(¨Ô/Ø ™(ˆŒØ38ˆ×ÑÔ0à/4ö8
Ø*+ˆA�×%Ñ%Ó%ò8
ˆ×ÑÔ4ð 49ö3
Ø./ˆA�×%Ñ%Ñ%¨Ó)ò3
ˆ×ÑÔ/ð )-¯©¸4×;OÑ;OÑ(OˆÕ%ùò8
ùò3
s   ÁCÁ=C) é   i   r*   rB  r   FFr   FFFr   r   FFFr   FFr*   TFFFFFFFFFNF©rj   Úboolro   rl  rp   rl  rq   rl  rr   rl  rt   rl  ru   rl  r}   rl  )NNNF)NN©rð   ztorch.Tensor)r   )F)FTN)r5   r6   r7   r8   r   r9   r   Úpropertyr�   r{   rv   rß   ré   rò   rþ   r@  rM  rT  rX  r[  rP  ÚstaticmethodrF  rh  r)   r:   r;   s   @r%   r=   r=   [   sz  ø… ñCðJ  €N�EÓð ØØØ"%Ø!$Ø+0Ø-2Ø "Ø"'Ø%*Ø %ØØØ&+Ø05Ø#(Ø()Ø&+Ø %Ø "Ø!%Ø"'Ø(-Ø+0Ø %Ø$)Ø!&Ø %Ø).Ø%*Ø$(Ø&+ðItð& *.ð'tð: ð;tð< "ð=tð> ð?tð@ ðAtðB #'ðCtðD #ðEtðH  $õItðl ñ.ó ð.ò
ò-9ðd ØØØóR
óh-
ó^óIð: ób!òH)?òV88òt2ðH ð à!ó':òx
ð ñó ðó";öPr&   r=   c                  ó°   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Zd	d„Zd
ˆ fd„Zedd„«       Zd„ Z	ˆ xZ
S )Ú	SAM3ModelúfSAM3Model class for Segment Anything Model 3 with memory-based video object segmentation capabilities.c$                ó¤  •— t        ‰$| �  g |‘|‘|‘|‘|‘|‘|‘|‘|	‘|
‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘|‘| ‘|!‘|"‘|#‘­Ž  t        d	dt        d| j                  dd¬«      | j                  dd| j
                  | j                  | j                  | j                  | j                  dœ
| j                  xs i ¤Ž| _        y)
rr  r?   rŠ   r‹   rŒ   r�   r’   r“   Nrš   )r   r   r   r
   r›   rI   rk   ro   rp   rj   rn   rž   r|   s%                                       €r%   r   zSAM3Model.__init__  sú  ø€ ôN 	‰Ñð $	
Øð$	
àð$	
ð ð$	
ð ð	$	
ð
 ð$	
ð ð$	
ð &ð$	
ð %ð$	
ð /ð$	
ð 1ð$	
ð $ð$	
ð &ð$	
ð )ð$	
ð $ð$	
ð !ð$	
ð  !ð!$	
ð" *ð#$	
ð$ ,ð%$	
ð& 'ð'$	
ð( ,ð)$	
ð* *ð+$	
ð, $ð-$	
ð. $ð/$	
ð0 %ð1$	
ð2 &ð3$	
ð4 ,ð5$	
ð6 /ð7$	
ð8 ð9$	
ð:  ð;$	
ð< ð=$	
ð> ð?$	
ð@ %ðA$	
ðB !ðC$	
ðD (ðE$	
ðF "óG$	
ôJ !0ð !
Ø"#Ü)ØØ"×7Ñ7ØØô	ð !×5Ñ5ØØ #Ø"&×"CÑ"CØ'+×'FÑ'FØ ×0Ñ0Ø $× 8Ñ 8Ø,0×,PÑ,Pñ!
ð  ×/Ñ/Ò5°2ñ!!
ˆÕr&   c                óð   — | j                   j                  |«      }| j                  rN| j                  j	                  |d   d   «      |d   d<   | j                  j                  |d   d   «      |d   d<   |S rë   )r   Úforward_image_sam2rI   rž   rí   rî   rï   s      r%   rò   zSAM3Model.forward_imagec  s�   € à×)Ñ)×<Ñ<¸YÓGˆØ×,Ò,ð /3×.CÑ.C×.KÑ.KÈLÐYgÑLhÐijÑLkÓ.lˆL˜Ñ(¨Ñ+Ø.2×.CÑ.C×.KÑ.KÈLÐYgÑLhÐijÑLkÓ.lˆL˜Ñ(¨Ñ+ØÐr&   c                óŒ   •— t         ‰| �  |«       |D �cg c]
  }|dz  dz  ‘Œ c}| j                  j                  _        yc c}w )z6Set the image size for the model and mask downsampler.é   r*   N)r   r)   rV   Úmask_downsamplerÚinterpol_size)r#   r0   r«   r$   s      €r%   r)   zSAM3Model.set_imgszm  s;   ø€ ä‰Ñ˜%Ô ØZ_Ö=`ÐRV¸dÀb¹jÈ2»oÒ=`ˆ×Ñ×,Ñ,Õ:ùÒ=`s   •Ac                ó  — | dkD  j                  d¬«      }|dkD  j                  d¬«      }t        j                  |d¬«      }||z  }||k\  }|d   j                  | «      }t        j                  || t        j                  | d¬«      «      }|S )	zXSuppress masks that shrink in area after applying pixelwise non-overlapping constraints.r   )r   r¨   r¸   rB  )r  rE  rá   r`  )Úsumr    rb  Ú	expand_asrÃ   )	rZ  Únew_pred_masksÚshrink_thresholdÚarea_beforeÚ
area_afterÚ
area_ratiorf  Ú	keep_maskÚpred_masks_afters	            r%   Ú_suppress_shrinked_masksz"SAM3Model._suppress_shrinked_masksr  s’   € ð " A‘~×*Ñ*¨xÐ*Ó8ˆØ$ qÑ(×-Ñ-°(Ð-Ó;ˆ
Ü—k‘k +°3Ô7ˆØ +Ñ-ˆ
ØÐ-Ñ-ˆØ˜Ñ)×3Ñ3°JÓ?ˆ	Ü Ÿ;™; y°*¼e¿k¹kÈ*ÐZ_Ô>`ÓaÐØÐr&   c                óL   — | j                  |«      }| j                  ||«      }|S )z®This function suppresses masks that shrink in area after applying pixelwise non-overlapping constraints. Note
        that the final output can still be overlapping.
        )rF  r„  )r#   rZ  Ú!pixel_level_non_overlapping_maskss      r%   Ú"_suppress_object_pw_area_shrinkagez,SAM3Model._suppress_object_pw_area_shrinkage~  s1   € ð
 -1×,SÑ,SÐT^Ó,_Ð)ð ×2Ñ2°:Ð?`Óaˆ
ØÐr&   ) rj  ið  rw  r   r   FFr   FFFr   r   FFFr   FFr*   TFFFFFFFFFNFrk  rm  )r0   ztuple[int, int])g333333Ó?)r5   r6   r7   r8   r   rò   r)   ro  r„  r‡  r:   r;   s   @r%   rq  rq    sô   ø„ Ùpð ØØØ"#Ø!"Ø+0Ø-2Ø "Ø"'Ø%*Ø %ØØØ&+Ø05Ø#(Ø()Ø&+Ø %Ø "Ø!%Ø"'Ø(-Ø+0Ø %Ø$)Ø!&Ø %Ø).Ø%*Ø$(Ø&+ðI]
ð& *.ð']
ð: ð;]
ð< "ð=]
ð> ð?]
ð@ ðA]
ðB #'ðC]
ðD #ðE]
ðH  $õI]
ó~õað
 ò	 ó ð	 ö	r&   rq  )Ú
__future__r   r    Útorch.nn.functionalr   Ú
functionalr¿   Útorch.nn.initr   Úultralytics.nn.modulesr   Úultralytics.utilsr   Úblocksr	   r
   Údecodersr   r   Úencodersr   r   Úutilsr   r   rÄ   ÚModuler   r=   rq  rš   r&   r%   ú<module>r“     sj   ðõ #ã ß Ð Ý Ý 'å &Ý $ç <ß 2ß 4ß =ð €ô?/ˆr�y‰yô ?/ôDcP�—‘—‘ô cPôLF�	õ Fr&   