Ë
    Fêñi¹  ã                   óf   — d dl Zd dlZd dlmZ d dlmZ d dlmZ  G d„ de«      Z	 G d„ de	e«      Z
y)	é    N)ÚLoadVisualPrompt)ÚDetectionPredictor)ÚSegmentationPredictorc                   óV   ‡ — e Zd ZdZd	defˆ fd„Zd„ Zˆ fd„Zd
ˆ fd„	Zˆ fd„Z	d„ Z
ˆ xZS )ÚYOLOEVPDetectPredictoraa  A class extending DetectionPredictor for YOLO-EVP (Enhanced Visual Prompting) predictions.

    This class provides common functionality for YOLO models that use visual prompting, including model setup, prompt
    handling, and preprocessing transformations.

    Attributes:
        model (torch.nn.Module): The YOLO model for inference.
        device (torch.device): Device to run the model on (CPU or CUDA).
        prompts (dict | torch.Tensor): Visual prompts containing class indices and bounding boxes or masks.

    Methods:
        setup_model: Initialize the YOLO model and set it to evaluation mode.
        set_prompts: Set the visual prompts for the model.
        pre_transform: Preprocess images and prompts before inference.
        inference: Run inference with visual prompts.
        get_vpe: Process source to get visual prompt embeddings.
    Úverbosec                 ó6   •— t         ‰| �  ||¬«       d| _        y)z½Set up the model for prediction.

        Args:
            model (torch.nn.Module): Model to load or use.
            verbose (bool, optional): If True, provides detailed logging.
        )r   TN)ÚsuperÚsetup_modelÚdone_warmup)ÚselfÚmodelr   Ú	__class__s      €úg/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/yolo/yoloe/predict.pyr   z"YOLOEVPDetectPredictor.setup_model   s   ø€ ô 	‰Ñ˜E¨7ÐÔ3ØˆÕó    c                 ó   — || _         y)z×Set the visual prompts for the model.

        Args:
            prompts (dict): Dictionary containing class indices and bounding boxes or masks. Must include a 'cls' key
                with class indices.
        N)Úprompts)r   r   s     r   Úset_promptsz"YOLOEVPDetectPredictor.set_prompts(   s   € ð ˆ�r   c           
      óä  •— t         ‰	| �  |«      }| j                  j                  dd«      }| j                  j                  dd«      }| j                  d   }t	        |«      dk(  ra| j                  |d   j                  dd |d   j                  dd |||«      }|j                  d«      j                  | j                  «      }�nb|€J d|› d	�«       ‚t        |t        «      rt        d
„ |D «       «      sJ d|› d	�«       ‚t        |t        «      rt        d„ |D «       «      sJ d|› d	�«       ‚t	        |«      t	        |«      cxk(  rt	        |«      k(  s.n J dt	        |«      › dt	        |«      › dt	        |«      › d	�«       ‚t        t	        |«      «      D �cg c]<  }| j                  ||   j                  dd ||   j                  dd ||   ||   «      ‘Œ> }}t        j                  j                   j"                  j%                  |d¬«      j                  | j                  «      }| j&                  j(                  r|j+                  «       | _        |S |j-                  «       | _        |S c c}w )aÇ  Preprocess images and prompts before inference.

        This method applies letterboxing to the input image and transforms the visual prompts (bounding boxes or masks)
        accordingly.

        Args:
            im (list): List of input images.

        Returns:
            (list): Preprocessed images ready for model inference.

        Raises:
            ValueError: If neither valid bounding boxes nor masks are provided in the prompts.
        ÚbboxesNÚmasksÚclsé   r   é   zExpected bboxes, but got ú!c              3   óP   K  — | ]  }t        |t        j                  «      –— Œ  y ­w©N©Ú
isinstanceÚnpÚndarray©Ú.0Úbs     r   ú	<genexpr>z7YOLOEVPDetectPredictor.pre_transform.<locals>.<genexpr>K   s   è ø€ Ò3^ÐRS´J¸qÄ"Ç*Á*×4MÑ3^ùó   ‚$&z#Expected list[np.ndarray], but got c              3   óP   K  — | ]  }t        |t        j                  «      –— Œ  y ­wr   r   r"   s     r   r%   z7YOLOEVPDetectPredictor.pre_transform.<locals>.<genexpr>N   s   è ø€ Ò5bÐTU´jÀÄBÇJÁJ×6OÑ5bùr&   z-Expected same length for all inputs, but got ÚvsT)Úbatch_first)r
   Úpre_transformr   ÚpopÚlenÚ_process_single_imageÚshapeÚ	unsqueezeÚtoÚdevicer   ÚlistÚallÚrangeÚtorchÚnnÚutilsÚrnnÚpad_sequencer   Úfp16ÚhalfÚfloat)
r   ÚimÚimgr   r   ÚcategoryÚvisualsr   Úir   s
            €r   r*   z$YOLOEVPDetectPredictor.pre_transform1   sR  ø€ ô ‰gÑ# BÓ'ˆØ—‘×!Ñ! (¨DÓ1ˆØ—‘× Ñ  ¨$Ó/ˆØ—<‘< Ñ&ˆÜˆs‹8�qŠ=Ø×0Ñ0°°Q±·±¸b¸qÐ1AÀ2ÀaÁ5Ç;Á;ÈrÐPQÀ?ÐT\Ð^dÐfkÓlˆGØ×'Ñ'¨Ó*×-Ñ-¨d¯k©kÓ:ŠGð Ð%ÐLÐ)BÀ6À(È!Ð'LÓLÐ%ä˜f¤dÔ+´Ñ3^ÐW]Ô3^Ô0^ð Ø5°f°X¸QÐ?óÐ^ô ˜h¬Ô-´#Ñ5bÐYaÔ5bÔ2bð Ø5°h°Z¸qÐAóÐbô �r“7œc (›mÔ:¬s°6«{Ô:ð Ø?ÄÀBÃ¸yÈÌ3ÈxË=È/ÐY[Ô\_Ð`fÓ\gÐ[hÐhiÐjóÐ:ô
 œs 3›x›öàð ×*Ñ*¨3¨q©6¯<©<¸¸Ð+;¸RÀ¹U¿[¹[ÈÈ!¸_ÈhÐWXÉkÐ[aÐbcÑ[dÕeðˆGð ô —h‘h—n‘n×(Ñ(×5Ñ5°gÈ4Ð5ÓP×SÑSÐTX×T_ÑT_Ó`ˆGØ)-¯©¯ª�w—|‘|“~ˆŒØˆ
ð ?F¿m¹m»oˆŒØˆ
ùòs   ÆAI-c           
      ód  •— |�Øt        |«      rÍt        j                  |t        j                  ¬«      }|j                  dk(  r	|ddd…f   }t        |d   |d   z  |d   |d   z  «      }||z  }|dddd…fxx   t        |d   t        |d   |z  «      z
  dz  dz
  «      z  cc<   |dddd…fxx   t        |d   t        |d   |z  «      z
  dz  dz
  «      z  cc<   n:|�-t        ‰| �!  |«      }t        j                  |«      }d||dk(  <   nt        d	«      ‚t        «       j                  ||||«      S )
aÆ  Process a single image by resizing bounding boxes or masks and generating visuals.

        Args:
            dst_shape (tuple): The target shape (height, width) of the image.
            src_shape (tuple): The original shape (height, width) of the image.
            category (list | np.ndarray): The category indices for visual prompts.
            bboxes (list | np.ndarray, optional): A list of bounding boxes in the format [x1, y1, x2, y2].
            masks (np.ndarray, optional): A list of masks corresponding to the image.

        Returns:
            (torch.Tensor): The processed visuals for the image.

        Raises:
            ValueError: If neither `bboxes` nor `masks` are provided.
        N)Údtyper   r   .r   gš™™™™™¹?ér   z$Please provide valid bboxes or masks)r,   r    ÚarrayÚfloat32ÚndimÚminÚroundr
   r*   ÚstackÚ
ValueErrorr   Úget_visuals)	r   Ú	dst_shapeÚ	src_shaper?   r   r   ÚgainÚresized_masksr   s	           €r   r-   z,YOLOEVPDetectPredictor._process_single_image\   sJ  ø€ ð  Ð¤# f¤+Ü—X‘X˜f¬B¯J©JÔ7ˆFØ�{‰{˜aÒØ ¢a ™�ä�y ‘| i°¡lÑ2°I¸a±LÀ9ÈQÁ<Ñ4OÓPˆDØ�d‰NˆFØ�3˜˜˜1˜�9Ó¤¨	°!©´u¸YÀq¹\ÈDÑ=PÓ7QÑ(QÐUVÑ'VÐY\Ñ'\Ó!]Ñ]ÓØ�3˜˜˜1˜�9Ó¤¨	°!©´u¸YÀq¹\ÈDÑ=PÓ7QÑ(QÐUVÑ'VÐY\Ñ'\Ó!]Ñ]ÔØÐä!™GÑ1°%Ó8ˆMÜ—H‘H˜]Ó+ˆEØ"#ˆE�%˜3‘,ÒäÐCÓDÐDô  Ó!×-Ñ-¨h¸	À6È5ÓQÐQr   c                 óB   •— t        ‰| �  |g|¢­d| j                  i|¤ŽS )a&  Run inference with visual prompts.

        Args:
            im (torch.Tensor): Input image tensor.
            *args (Any): Variable length argument list.
            **kwargs (Any): Arbitrary keyword arguments.

        Returns:
            (torch.Tensor): Model prediction results.
        Úvpe)r
   Ú	inferencer   )r   r=   ÚargsÚkwargsr   s       €r   rS   z YOLOEVPDetectPredictor.inference€   s(   ø€ ô ‰wÑ  ÐG¸ÒG¨¯©ÐGÀÑGÐGr   c                 óî   — | j                  |«       t        | j                  «      dk(  sJ d«       ‚| j                  D ]6  \  }}}| j                  |«      }| j	                  || j
                  d¬«      c S  y)aÃ  Process the source to get the visual prompt embeddings (VPE).

        Args:
            source (str | Path | int | PIL.Image | np.ndarray | torch.Tensor | list | tuple): The source of the image to
                make predictions on. Accepts various types including file paths, URLs, PIL images, numpy arrays, and
                torch tensors.

        Returns:
            (torch.Tensor): The visual prompt embeddings (VPE) from the model.
        r   z get_vpe only supports one image!T)rR   Ú
return_vpeN)Úsetup_sourcer,   ÚdatasetÚ
preprocessr   r   )r   ÚsourceÚ_Úim0sr=   s        r   Úget_vpezYOLOEVPDetectPredictor.get_vpe�   sq   € ð 	×Ñ˜&Ô!Ü�4—<‘<Ó  AÒ%ÐIÐ'IÓIÐ%ØŸ,™,ò 	E‰JˆAˆt�QØ—‘ Ó&ˆBØ—:‘:˜b d§l¡l¸t�:ÓDÒDñ	Er   )T)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úboolr   r   r*   r-   rS   r^   Ú__classcell__)r   s   @r   r   r      s2   ø„ ññ$ ¨$õ  òô)õV"RôHHöEr   r   c                   ó   — e Zd ZdZy)ÚYOLOEVPSegPredictorz\Predictor for YOLO-EVP segmentation tasks combining detection and segmentation capabilities.N)r_   r`   ra   rb   © r   r   rf   rf   Ÿ   s   „ Ùfàr   rf   )Únumpyr    r5   Úultralytics.data.augmentr   Úultralytics.models.yolo.detectr   Úultralytics.models.yolo.segmentr   r   rf   rg   r   r   ú<module>rl      s8   ðó Û å 5Ý =Ý AôQEÐ/ô QEôh	Ð0Ð2Gõ 	r   