Ë
    FêñiÕa  ã                  óˆ   — d dl mZ d dlZd dlmZ d dlmZmZ  G d„ dej                  «      Z G d„ dej                  «      Z	y)	é    )ÚannotationsN)Únn)ÚMLPÚLayerNorm2dc                  ó˜   ‡ — e Zd ZdZdej
                  ddf	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Z	 	 	 	 	 	 	 	 	 	 	 	 dd„Z	 	 	 	 	 	 	 	 	 	 d	d„Zˆ xZ	S )
ÚMaskDecoderaà  Decoder module for generating masks and their associated quality scores using a transformer architecture.

    This class predicts masks given image and prompt embeddings, utilizing a transformer to process the inputs and
    generate mask predictions along with their quality scores.

    Attributes:
        transformer_dim (int): Channel dimension for the transformer module.
        transformer (nn.Module): Transformer module used for mask prediction.
        num_multimask_outputs (int): Number of masks to predict for disambiguating masks.
        iou_token (nn.Embedding): Embedding for the IoU token.
        num_mask_tokens (int): Number of mask tokens.
        mask_tokens (nn.Embedding): Embedding for the mask tokens.
        output_upscaling (nn.Sequential): Neural network sequence for upscaling the output.
        output_hypernetworks_mlps (nn.ModuleList): Hypernetwork MLPs for generating masks.
        iou_prediction_head (nn.Module): MLP for predicting mask quality.

    Methods:
        forward: Predict masks given image and prompt embeddings.
        predict_masks: Internal method for mask prediction.

    Examples:
        >>> decoder = MaskDecoder(transformer_dim=256, transformer=transformer_module)
        >>> masks, iou_pred = decoder(
        ...     image_embeddings, image_pe, sparse_prompt_embeddings, dense_prompt_embeddings, multimask_output=True
        ... )
        >>> print(f"Predicted masks shape: {masks.shape}, IoU predictions shape: {iou_pred.shape}")
    é   é   c                óŽ  •— t         ‰| �  «        || _        || _        || _        t        j                  d|«      | _        |dz   | _        t        j                  | j                  |«      | _	        t        j                  t        j                  ||dz  dd¬«      t        |dz  «       |«       t        j                  |dz  |dz  dd¬«       |«       «      | _        t        j                  t        | j                  «      D �cg c]  }t!        |||dz  d«      ‘Œ c}«      | _        t!        ||| j                  |«      | _        yc c}w )a€  Initialize the MaskDecoder module for generating masks and their associated quality scores.

        Args:
            transformer_dim (int): Channel dimension for the transformer module.
            transformer (nn.Module): Transformer module used for mask prediction.
            num_multimask_outputs (int): Number of masks to predict for disambiguating masks.
            activation (type[nn.Module]): Type of activation to use when upscaling masks.
            iou_head_depth (int): Depth of the MLP used to predict mask quality.
            iou_head_hidden_dim (int): Hidden dimension of the MLP used to predict mask quality.
        é   é   é   ©Úkernel_sizeÚstrideé   r	   N)ÚsuperÚ__init__Útransformer_dimÚtransformerÚnum_multimask_outputsr   Ú	EmbeddingÚ	iou_tokenÚnum_mask_tokensÚmask_tokensÚ
SequentialÚConvTranspose2dr   Úoutput_upscalingÚ
ModuleListÚranger   Úoutput_hypernetworks_mlpsÚiou_prediction_head)	Úselfr   r   r   Ú
activationÚiou_head_depthÚiou_head_hidden_dimÚ_Ú	__class__s	           €úi/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/sam/modules/decoders.pyr   zMaskDecoder.__init__(   s$  ø€ ô& 	‰ÑÔØ.ˆÔØ&ˆÔà%:ˆÔ"äŸ™ a¨Ó9ˆŒØ4°qÑ8ˆÔÜŸ<™<¨×(<Ñ(<¸oÓNˆÔä "§¡Ü×Ñ˜°À1Ñ0DÐRSÐ\]Ô^Ü˜¨1Ñ,Ó-Ù‹LÜ×Ñ˜°!Ñ3°_ÈÑ5IÐWXÐabÔcÙ‹Ló!
ˆÔô *,¯©ÜUZÐ[_×[oÑ[oÓUpÖqÐPQŒS� /°?ÀaÑ3GÈÕKÒqó*
ˆÔ&ô $' Ð8KÈT×MaÑMaÐcqÓ#rˆÕ ùò rs   Ä Ec                óž   — | j                  ||||¬«      \  }}|rt        dd«      nt        dd«      }|dd…|dd…dd…f   }|dd…|f   }||fS )a   Predict masks given image and prompt embeddings.

        Args:
            image_embeddings (torch.Tensor): Embeddings from the image encoder.
            image_pe (torch.Tensor): Positional encoding with the shape of image_embeddings.
            sparse_prompt_embeddings (torch.Tensor): Embeddings of the points and boxes.
            dense_prompt_embeddings (torch.Tensor): Embeddings of the mask inputs.
            multimask_output (bool): Whether to return multiple masks or a single mask.

        Returns:
            masks (torch.Tensor): Batched predicted masks.
            iou_pred (torch.Tensor): Batched predictions of mask quality.

        Examples:
            >>> decoder = MaskDecoder(transformer_dim=256, transformer=transformer_module)
            >>> image_emb = torch.rand(1, 256, 64, 64)
            >>> image_pe = torch.rand(1, 256, 64, 64)
            >>> sparse_emb = torch.rand(1, 2, 256)
            >>> dense_emb = torch.rand(1, 256, 64, 64)
            >>> masks, iou_pred = decoder(image_emb, image_pe, sparse_emb, dense_emb, multimask_output=True)
            >>> print(f"Masks shape: {masks.shape}, IoU predictions shape: {iou_pred.shape}")
        )Úimage_embeddingsÚimage_peÚsparse_prompt_embeddingsÚdense_prompt_embeddingsr   Nr   )Úpredict_masksÚslice)	r#   r+   r,   r-   r.   Úmultimask_outputÚmasksÚiou_predÚ
mask_slices	            r)   ÚforwardzMaskDecoder.forwardR   sk   € ð< ×,Ñ,Ø-ØØ%=Ø$;ð	 -ó 
‰ˆˆxñ (8”U˜1˜d”^¼UÀ1Àa»[ˆ
Ø’a˜¢QªÐ)Ñ*ˆØšA˜z˜MÑ*ˆà�hˆÐó    c           
     ó  — t        j                  | j                  j                  | j                  j                  gd¬«      }|j                  d«      j                  |j                  d   dd«      }t        j                  ||fd¬«      }t        j                  ||j                  d   d¬«      }||z   }t        j                  ||j                  d   d¬«      }|j                  \  }	}
}}| j                  |||«      \  }}|dd…ddd…f   }|dd…dd| j                  z   …dd…f   }|j                  dd«      j                  |	|
||«      }| j                  |«      }t        | j                  «      D �cg c]!  } | j                  |   |dd…|dd…f   «      ‘Œ# }}t        j                   |d¬«      }|j                  \  }	}
}}||j                  |	|
||z  «      z  j                  |	d||«      }| j#                  |«      }||fS c c}w )z`Predict masks and quality scores using image and prompt embeddings via transformer architecture.r   ©Údiméÿÿÿÿr   Nr   )ÚtorchÚcatr   Úweightr   Ú	unsqueezeÚexpandÚshapeÚrepeat_interleaver   r   Ú	transposeÚviewr   r    r!   Ústackr"   )r#   r+   r,   r-   r.   Úoutput_tokensÚtokensÚsrcÚpos_srcÚbÚcÚhÚwÚhsÚiou_token_outÚmask_tokens_outÚupscaled_embeddingÚiÚhyper_in_listÚhyper_inr2   r3   s                         r)   r/   zMaskDecoder.predict_masks~   s  € ô Ÿ	™	 4§>¡>×#8Ñ#8¸$×:JÑ:J×:QÑ:QÐ"RÐXYÔZˆØ%×/Ñ/°Ó2×9Ñ9Ð:R×:XÑ:XÐYZÑ:[Ð]_ÐacÓdˆÜ—‘˜MÐ+CÐDÈ!ÔLˆô ×%Ñ%Ð&6¸¿¹ÀQ¹ÈQÔOˆØÐ+Ñ+ˆÜ×)Ñ)¨(°F·L±LÀ±OÈÔKˆØ—Y‘Y‰
ˆˆ1ˆa�ð ×"Ñ" 3¨°Ó8‰ˆˆCØš1˜a¢˜7™ˆØšQ  Q¨×)=Ñ)=Ñ%=Ð >ÂÐAÑBˆð �m‰m˜A˜qÓ!×&Ñ& q¨!¨Q°Ó2ˆØ!×2Ñ2°3Ó7ÐäQVÐW[×WkÑWkÓQlö-
ØLMÐ-ˆD×*Ñ*¨1Ñ-¨oºaÀÂA¸gÑ.FÕGð-
ˆð -
ô —;‘;˜}°!Ô4ˆØ'×-Ñ-‰
ˆˆ1ˆa�ØÐ.×3Ñ3°A°q¸!¸a¹%Ó@Ñ@×FÑFÀqÈ"ÈaÐQRÓSˆð ×+Ñ+¨MÓ:ˆà�hˆÐùò-
s   Å3&H)r   Úintr   ú	nn.Moduler   rT   r$   útype[nn.Module]r%   rT   r&   rT   ÚreturnÚNone)r+   útorch.Tensorr,   rY   r-   rY   r.   rY   r1   ÚboolrW   ú!tuple[torch.Tensor, torch.Tensor])
r+   rY   r,   rY   r-   rY   r.   rY   rW   r[   )
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚGELUr   r5   r/   Ú__classcell__©r(   s   @r)   r   r      sâ   ø„ ñð@ &'Ø&(§g¡gØØ#&ð(sàð(sð ð(sð  #ð	(sð
 $ð(sð ð(sð !ð(sð 
õ(sðT*à&ð*ð ð*ð #/ð	*ð
 ".ð*ð ð*ð 
+ó*ðX%à&ð%ð ð%ð #/ð	%ð
 ".ð%ð 
+÷%r6   r   c                  óØ   ‡ — e Zd ZdZdej
                  ddddddddddf	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Z	 d	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd„Z	 d	 	 	 	 	 	 	 	 	 	 	 	 	 dd	„Zd
„ Z	d„ Z
ˆ xZS )ÚSAM2MaskDecodera~
  Transformer-based decoder for predicting instance segmentation masks from image and prompt embeddings.

    This class extends the functionality of the MaskDecoder, incorporating additional features such as high-resolution
    feature processing, dynamic multimask output, and object score prediction.

    Attributes:
        transformer_dim (int): Channel dimension of the transformer.
        transformer (nn.Module): Transformer used to predict masks.
        num_multimask_outputs (int): Number of masks to predict when disambiguating masks.
        iou_token (nn.Embedding): Embedding for IOU token.
        num_mask_tokens (int): Total number of mask tokens.
        mask_tokens (nn.Embedding): Embedding for mask tokens.
        pred_obj_scores (bool): Whether to predict object scores.
        obj_score_token (nn.Embedding): Embedding for object score token.
        use_multimask_token_for_obj_ptr (bool): Whether to use multimask token for object pointer.
        output_upscaling (nn.Sequential): Upscaling layers for output.
        use_high_res_features (bool): Whether to use high-resolution features.
        conv_s0 (nn.Conv2d): Convolutional layer for high-resolution features (s0).
        conv_s1 (nn.Conv2d): Convolutional layer for high-resolution features (s1).
        output_hypernetworks_mlps (nn.ModuleList): List of MLPs for output hypernetworks.
        iou_prediction_head (MLP): MLP for IOU prediction.
        pred_obj_score_head (nn.Linear | MLP): Linear layer or MLP for object score prediction.
        dynamic_multimask_via_stability (bool): Whether to use dynamic multimask via stability.
        dynamic_multimask_stability_delta (float): Delta value for dynamic multimask stability.
        dynamic_multimask_stability_thresh (float): Threshold for dynamic multimask stability.

    Methods:
        forward: Predict masks given image and prompt embeddings.
        predict_masks: Predict instance segmentation masks from image and prompt embeddings.
        _get_stability_scores: Compute mask stability scores based on IoU between thresholds.
        _dynamic_multimask_via_stability: Dynamically select the most stable mask output.

    Examples:
        >>> image_embeddings = torch.rand(1, 256, 64, 64)
        >>> image_pe = torch.rand(1, 256, 64, 64)
        >>> sparse_prompt_embeddings = torch.rand(1, 2, 256)
        >>> dense_prompt_embeddings = torch.rand(1, 256, 64, 64)
        >>> decoder = SAM2MaskDecoder(256, transformer)
        >>> masks, iou_pred, sam_tokens_out, obj_score_logits = decoder.forward(
        ...     image_embeddings, image_pe, sparse_prompt_embeddings, dense_prompt_embeddings, True, False
        ... )
    r	   r
   Fgš™™™™™©?g\�Âõ(\ï?c                ó4  •— t         ‰| �  «        || _        || _        || _        t        j                  d|«      | _        |dz   | _        t        j                  | j                  |«      | _	        || _
        | j                  rt        j                  d|«      | _        || _        t        j                  t        j                  ||dz  dd¬«      t        |dz  «       |«       t        j                  |dz  |dz  dd¬«       |«       «      | _        || _        |rBt        j$                  ||dz  dd¬«      | _        t        j$                  ||dz  dd¬«      | _        t        j*                  t-        | j                  «      D �cg c]  }t/        |||dz  d«      ‘Œ c}«      | _        t/        ||| j                  ||¬«      | _        | j                  r0t        j4                  |d«      | _        |rt/        ||dd«      | _        |	| _        |
| _        || _        yc c}w )	a  Initialize the SAM2MaskDecoder module for predicting instance segmentation masks.

        This decoder extends the functionality of MaskDecoder, incorporating additional features such as high-resolution
        feature processing, dynamic multimask output, and object score prediction.

        Args:
            transformer_dim (int): Channel dimension of the transformer.
            transformer (nn.Module): Transformer used to predict masks.
            num_multimask_outputs (int): Number of masks to predict when disambiguating masks.
            activation (type[nn.Module]): Type of activation to use when upscaling masks.
            iou_head_depth (int): Depth of the MLP used to predict mask quality.
            iou_head_hidden_dim (int): Hidden dimension of the MLP used to predict mask quality.
            use_high_res_features (bool): Whether to use high-resolution features.
            iou_prediction_use_sigmoid (bool): Whether to use sigmoid for IOU prediction.
            dynamic_multimask_via_stability (bool): Whether to use dynamic multimask via stability.
            dynamic_multimask_stability_delta (float): Delta value for dynamic multimask stability.
            dynamic_multimask_stability_thresh (float): Threshold for dynamic multimask stability.
            pred_obj_scores (bool): Whether to predict object scores.
            pred_obj_scores_mlp (bool): Whether to use MLP for object score prediction.
            use_multimask_token_for_obj_ptr (bool): Whether to use multimask token for object pointer.
        r   r   r   r   r   r	   )ÚsigmoidN)r   r   r   r   r   r   r   r   r   r   Úpred_obj_scoresÚobj_score_tokenÚuse_multimask_token_for_obj_ptrr   r   r   r   Úuse_high_res_featuresÚConv2dÚconv_s0Úconv_s1r   r    r   r!   r"   ÚLinearÚpred_obj_score_headÚdynamic_multimask_via_stabilityÚ!dynamic_multimask_stability_deltaÚ"dynamic_multimask_stability_thresh)r#   r   r   r   r$   r%   r&   rj   Úiou_prediction_use_sigmoidrp   rq   rr   rg   Úpred_obj_scores_mlpri   r'   r(   s                   €r)   r   zSAM2MaskDecoder.__init__Ò   sî  ø€ ôL 	‰ÑÔØ.ˆÔØ&ˆÔà%:ˆÔ"äŸ™ a¨Ó9ˆŒØ4°qÑ8ˆÔÜŸ<™<¨×(<Ñ(<¸oÓNˆÔà.ˆÔØ×ÒÜ#%§<¡<°°?Ó#CˆDÔ Ø/NˆÔ,ä "§¡Ü×Ñ˜°À1Ñ0DÐRSÐ\]Ô^Ü˜¨1Ñ,Ó-Ù‹LÜ×Ñ˜°!Ñ3°_ÈÑ5IÐWXÐabÔcÙ‹Ló!
ˆÔð &;ˆÔ"Ù ÜŸ9™9 _°oÈÑ6JÐXYÐbcÔdˆDŒLÜŸ9™9 _°oÈÑ6JÐXYÐbcÔdˆDŒLä)+¯©ÜUZÐ[_×[oÑ[oÓUpÖqÐPQŒS� /°?ÀaÑ3GÈÕKÒqó*
ˆÔ&ô $'ØØØ× Ñ ØØ.ô$
ˆÔ ð ×ÒÜ')§y¡y°À!Ó'DˆDÔ$Ù"Ü+.¨ÀÐQRÐTUÓ+V�Ô(ð 0OˆÔ,Ø1RˆÔ.Ø2TˆÕ/ùò' rs   Æ Hc                ób  — | j                  ||||||¬«      \  }}	}
}|r|dd…dd…dd…dd…f   }|	dd…dd…f   }	nJ| j                  r"| j                  s| j                  ||	«      \  }}	n|dd…dd…dd…dd…f   }|	dd…dd…f   }	|r| j                  r|
dd…dd…f   }n|
dd…dd…f   }||	||fS )a›  Predict masks given image and prompt embeddings.

        Args:
            image_embeddings (torch.Tensor): Embeddings from the image encoder with shape (B, C, H, W).
            image_pe (torch.Tensor): Positional encoding with the shape of image_embeddings (B, C, H, W).
            sparse_prompt_embeddings (torch.Tensor): Embeddings of the points and boxes with shape (B, N, C).
            dense_prompt_embeddings (torch.Tensor): Embeddings of the mask inputs with shape (B, C, H, W).
            multimask_output (bool): Whether to return multiple masks or a single mask.
            repeat_image (bool): Flag to repeat the image embeddings.
            high_res_features (list[torch.Tensor] | None, optional): Optional high-resolution features.

        Returns:
            masks (torch.Tensor): Batched predicted masks with shape (B, N, H, W).
            iou_pred (torch.Tensor): Batched predictions of mask quality with shape (B, N).
            sam_tokens_out (torch.Tensor): Batched SAM token for mask output with shape (B, N, C).
            object_score_logits (torch.Tensor): Batched object score logits with shape (B, 1).

        Examples:
            >>> image_embeddings = torch.rand(1, 256, 64, 64)
            >>> image_pe = torch.rand(1, 256, 64, 64)
            >>> sparse_prompt_embeddings = torch.rand(1, 2, 256)
            >>> dense_prompt_embeddings = torch.rand(1, 256, 64, 64)
            >>> decoder = SAM2MaskDecoder(256, transformer)
            >>> masks, iou_pred, sam_tokens_out, obj_score_logits = decoder.forward(
            ...     image_embeddings, image_pe, sparse_prompt_embeddings, dense_prompt_embeddings, True, False
            ... )
        )r+   r,   r-   r.   Úrepeat_imageÚhigh_res_featuresNr   r   )r/   rp   ÚtrainingÚ _dynamic_multimask_via_stabilityri   )r#   r+   r,   r-   r.   r1   rv   rw   r2   r3   rO   Úobject_score_logitsÚsam_tokens_outs                r)   r5   zSAM2MaskDecoder.forward)  sï   € ðJ AE×@RÑ@RØ-ØØ%=Ø$;Ø%Ø/ð ASó A
Ñ=ˆˆx˜Ð*=ñ Øš!˜Q™R¢¢A˜+Ñ&ˆEØ¢ 1¡2 ‘‰HØ×1Ò1¸$¿-º-Ø"×CÑCÀEÈ8ÓT‰OˆE‘8àš!˜Q˜q˜S¢!¢Q˜,Ñ'ˆEØ¢ 1 Q 3 Ñ'ˆHá × DÒ DØ,ªQ°±¨UÑ3‰Nð -ªQ°°!°¨VÑ4ˆNà�h Ð0CÐCÐCr6   c           
     óª  — d}| j                   rYt        j                  | j                  j                  | j
                  j                  | j                  j                  gd¬«      }d}nAt        j                  | j
                  j                  | j                  j                  gd¬«      }|j                  d«      j                  |j                  d   dd«      }t        j                  ||fd¬«      }	|r&t        j                  ||	j                  d   d¬«      }
n#|j                  d   |	j                  d   k(  sJ ‚|}
|
|z   }
|j                  d   dk(  sJ d«       ‚t        j                  ||	j                  d   d¬«      }|
j                  \  }}}}| j                  |
||	«      \  }}
|dd…|dd…f   }|dd…|dz   |dz   | j                  z   …dd…f   }|
j                  dd«      j                  ||||«      }
| j                  r|€| j!                  |
«      }n?| j                   \  }}}}}|\  }} | | ||
«      |z   «      «      } | ||«      |z   «      }t#        | j                  «      D �cg c]!  } | j$                  |   |dd…|dd…f   «      ‘Œ# }}t        j&                  |d¬«      }|j                  \  }}}}||j                  ||||z  «      z  j                  |d||«      }| j)                  |«      }| j                   r#|dk(  sJ ‚| j+                  |dd…ddd…f   «      } n"d|j-                  |j                  d   d«      z  } |||| fS c c}w )	zYPredict instance segmentation masks from image and prompt embeddings using a transformer.r   r8   r   r:   z@image_pe should have size 1 in batch dim (from `get_dense_pe()`)Nr   g      $@)rg   r;   r<   rh   r=   r   r   r>   r?   r@   rA   r   r   rB   rC   rj   r   r    r!   rD   r"   ro   Únew_ones)!r#   r+   r,   r-   r.   rv   rw   ÚsrE   rF   rG   rH   rI   rJ   rK   rL   rM   rN   rO   rP   Údc1Úln1Úact1Údc2Úact2Úfeat_s0Úfeat_s1rQ   rR   rS   r2   r3   rz   s!                                    r)   r/   zSAM2MaskDecoder.predict_masksm  sd  € ð ˆØ×ÒÜ!ŸI™Ià×(Ñ(×/Ñ/Ø—N‘N×)Ñ)Ø×$Ñ$×+Ñ+ðð
 ôˆMð ‰Aä!ŸI™I t§~¡~×'<Ñ'<¸d×>NÑ>N×>UÑ>UÐ&VÐ\]Ô^ˆMØ%×/Ñ/°Ó2×9Ñ9Ð:R×:XÑ:XÐYZÑ:[Ð]_ÐacÓdˆÜ—‘˜MÐ+CÐDÈ!ÔLˆñ Ü×)Ñ)Ð*:¸F¿L¹LÈ¹OÐQRÔS‰Cà#×)Ñ)¨!Ñ,°·±¸Q±Ò?Ð?Ð?Ø"ˆCØÐ+Ñ+ˆØ�~‰~˜aÑ  AÒ%ÐiÐ'iÓiÐ%Ü×)Ñ)¨(°F·L±LÀ±OÈÔKˆØ—Y‘Y‰
ˆˆ1ˆa�ð ×"Ñ" 3¨°Ó8‰ˆˆCØš1˜a¢˜7™ˆØšQ  A¡¨¨Q©°×1EÑ1EÑ)EÐ FÊÐIÑJˆð �m‰m˜A˜qÓ!×&Ñ& q¨!¨Q°Ó2ˆØ×)Ò)Ð->Ð-FØ!%×!6Ñ!6°sÓ!;Ñà(,×(=Ñ(=Ñ%ˆC��d˜C Ø0ÑˆG�WÙ!%¡c©#¨c«(°WÑ*<Ó&=Ó!>ÐÙ!%¡cÐ*<Ó&=ÀÑ&GÓ!HÐô RWÐW[×WkÑWkÓQlö-
ØLMÐ-ˆD×*Ñ*¨1Ñ-¨oºaÀÂA¸gÑ.FÕGð-
ˆð -
ô —;‘;˜}°!Ô4ˆØ'×-Ñ-‰
ˆˆ1ˆa�ØÐ.×3Ñ3°A°q¸!¸a¹%Ó@Ñ@×FÑFÀqÈ"ÈaÐQRÓSˆð ×+Ñ+¨MÓ:ˆØ×ÒØ˜’6ˆM�6Ø"&×":Ñ":¸2ºaÀÂA¸g¹;Ó"GÑð #'¨×):Ñ):¸8¿>¹>È!Ñ;LÈaÓ)PÑ"PÐà�h Ð1DÐDÐDùò!-
s   É-&Mc                ó(  — |j                  d«      }t        j                  || j                  kD  d¬«      j	                  «       }t        j                  || j                   kD  d¬«      j	                  «       }t        j
                  |dkD  ||z  d«      S )zNCompute mask stability scores based on IoU between upper and lower thresholds.éþÿÿÿr:   r8   r   g      ð?)Úflattenr;   Úsumrq   ÚfloatÚwhere)r#   Úmask_logitsÚarea_iÚarea_us       r)   Ú_get_stability_scoresz%SAM2MaskDecoder._get_stability_scores´  sz   € à!×)Ñ)¨"Ó-ˆÜ—‘˜;¨×)OÑ)OÑOÐUWÔX×^Ñ^Ó`ˆÜ—‘˜;¨$×*PÑ*PÐ)PÑPÐVXÔY×_Ñ_ÓaˆÜ�{‰{˜6 A™: v°¡¸Ó<Ð<r6   c                óB  — |dd…dd…dd…dd…f   }|dd…dd…f   }t        j                  |d¬«      }t        j                  |j                  d   |j                  ¬«      }|||f   }|j                  d«      }|||f   }|j                  d«      }|dd…dd…dd…dd…f   }	|dd…dd…f   }
| j                  |	«      }|| j                  k\  }t        j                  |d   j                  |	«      |	|«      }t        j                  |j                  |
«      |
|«      }||fS )aƒ  Dynamically select the most stable mask output based on stability scores and IoU predictions.

        This method is used when outputting a single mask. If the stability score from the current single-mask output
        (based on output token 0) falls below a threshold, it instead selects from multi-mask outputs (based on output
        tokens 1-3) the mask with the highest predicted IoU score. This ensures a valid mask for both clicking and
        tracking scenarios.

        Args:
            all_mask_logits (torch.Tensor): Logits for all predicted masks, shape (B, N, H, W) where B is batch size, N
                is number of masks (typically 4), and H, W are mask dimensions.
            all_iou_scores (torch.Tensor): Predicted IoU scores for all masks, shape (B, N).

        Returns:
            mask_logits_out (torch.Tensor): Selected mask logits, shape (B, 1, H, W).
            iou_scores_out (torch.Tensor): Selected IoU scores, shape (B, 1).

        Examples:
            >>> decoder = SAM2MaskDecoder(...)
            >>> all_mask_logits = torch.rand(2, 4, 256, 256)  # 2 images, 4 masks each
            >>> all_iou_scores = torch.rand(2, 4)
            >>> mask_logits, iou_scores = decoder._dynamic_multimask_via_stability(all_mask_logits, all_iou_scores)
            >>> print(mask_logits.shape, iou_scores.shape)
            torch.Size([2, 1, 256, 256]) torch.Size([2, 1])
        Nr   r:   r8   r   )Údevice).NN)
r;   ÚargmaxÚaranger@   r‘   r>   r�   rr   r‹   Ú	expand_as)r#   Úall_mask_logitsÚall_iou_scoresÚmultimask_logitsÚmultimask_iou_scoresÚbest_scores_indsÚ
batch_indsÚbest_multimask_logitsÚbest_multimask_iou_scoresÚsinglemask_logitsÚsinglemask_iou_scoresÚstability_scoresÚ	is_stableÚmask_logits_outÚiou_scores_outs                  r)   ry   z0SAM2MaskDecoder._dynamic_multimask_via_stability»  sL  € ð4 +ª1¨a©b²!²Q¨;Ñ7ÐØ-ªa°±¨eÑ4ÐÜ Ÿ<™<Ð(<À"ÔEÐÜ—\‘\Ð"6×"<Ñ"<¸QÑ"?È×H]ÑH]Ô^ˆ
Ø 0°Ð=MÐ1MÑ NÐØ 5× ?Ñ ?ÀÓ BÐØ$8¸ÐEUÐ9UÑ$VÐ!Ø$=×$GÑ$GÈÓ$JÐ!ð ,ªA¨q°¨s²A²q¨LÑ9ÐØ .ªq°!°A°#¨vÑ 6ÐØ×5Ñ5Ð6GÓHÐØ$¨×(OÑ(OÑOˆ	ô  Ÿ+™+Ø�oÑ&×0Ñ0Ð1BÓCØØ!ó
ˆô
 Ÿ™Ø×ÑÐ 5Ó6Ø!Ø%ó
ˆð
  Ð.Ð.r6   )r   rT   r   rU   r   rT   r$   rV   r%   rT   r&   rT   rj   rZ   rg   rZ   rt   rZ   ri   rZ   rW   rX   )N)r+   rY   r,   rY   r-   rY   r.   rY   r1   rZ   rv   rZ   rw   úlist[torch.Tensor] | NonerW   ú=tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor])r+   rY   r,   rY   r-   rY   r.   rY   rv   rZ   rw   r£   rW   r¤   )r\   r]   r^   r_   r   r`   r   r5   r/   r�   ry   ra   rb   s   @r)   rd   rd   ¦   s“  ø„ ñ)ð^ &'Ø&(§g¡gØØ#&Ø&+Ø#(Ø(-Ø*.Ø+/Ø %Ø$)Ø05ðUUàðUUð ðUUð  #ð	UUð
 $ðUUð ðUUð !ðUUð  $ðUUð ðUUð "ðUUð *.ðUUð  
õ!UUð~ 8<ðBDà&ðBDð ðBDð #/ð	BDð
 ".ðBDð ðBDð ðBDð 5ðBDð 
GóBDðV 8<ðEEà&ðEEð ðEEð #/ð	EEð
 ".ðEEð ðEEð 5ðEEð 
GóEEòN=ö4/r6   rd   )
Ú
__future__r   r;   r   Úultralytics.nn.modulesr   r   ÚModuler   rd   © r6   r)   ú<module>r©      s8   ðõ #ã Ý ç 3ôX�"—)‘)ô XôvI/�b—i‘iõ I/r6   