Ë
    Fêñi{4  ã                  ó˜   — d dl mZ d dlmZ d dlmZ d dlZd dlZd dl	m
c mZ d dlmZ d dlmZmZ d dlmZ d dlmZmZ  G d	„ d
e«      Zy)é    )Úannotations)ÚPath)ÚAnyN)ÚDetectionValidator)ÚLOGGERÚops)Úcheck_requirements)ÚSegmentMetricsÚmask_iouc                  ó®   ‡ — e Zd ZdZddˆ fd„Zdˆ fd„Zdˆ fd„Zdd„Zdˆ fd„Zdˆ fd„Z	dˆ fd„Z
dˆ fd	„Zdˆ fd
„Zdd„Zdˆ fd„Zdˆ fd„Zdˆ fd„Zˆ xZS )ÚSegmentationValidatora–  A class extending the DetectionValidator class for validation based on a segmentation model.

    This validator handles the evaluation of segmentation models, processing both bounding box and mask predictions to
    compute metrics such as mAP for both detection and segmentation tasks.

    Attributes:
        plot_masks (list): List to store masks for plotting.
        process (callable): Function to process masks based on save_json and save_txt flags.
        args (SimpleNamespace): Arguments for the validator.
        metrics (SegmentMetrics): Metrics calculator for segmentation tasks.
        stats (dict): Dictionary to store statistics during validation.

    Examples:
        >>> from ultralytics.models.yolo.segment import SegmentationValidator
        >>> args = dict(model="yolo26n-seg.pt", data="coco8-seg.yaml")
        >>> validator = SegmentationValidator(args=args)
        >>> validator()
    c                óx   •— t         ‰| �  ||||«       d| _        d| j                  _        t        «       | _        y)a�  Initialize SegmentationValidator and set task to 'segment', metrics to SegmentMetrics.

        Args:
            dataloader (torch.utils.data.DataLoader, optional): DataLoader to use for validation.
            save_dir (Path, optional): Directory to save results.
            args (dict, optional): Arguments for the validator.
            _callbacks (dict, optional): Dictionary of callback functions.
        NÚsegment)ÚsuperÚ__init__ÚprocessÚargsÚtaskr
   Úmetrics)ÚselfÚ
dataloaderÚsave_dirr   Ú
_callbacksÚ	__class__s        €úe/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/yolo/segment/val.pyr   zSegmentationValidator.__init__&   s4   ø€ ô 	‰Ñ˜ X¨t°ZÔ@ØˆŒØ"ˆ�	‰	ŒÜ%Ó'ˆ�ó    c                óR   •— t         ‰| �  |«      }|d   j                  «       |d<   |S )zåPreprocess batch of images for YOLO segmentation validation.

        Args:
            batch (dict[str, Any]): Batch containing images and annotations.

        Returns:
            (dict[str, Any]): Preprocessed batch.
        Úmasks)r   Ú
preprocessÚfloat)r   Úbatchr   s     €r   r   z SegmentationValidator.preprocess4   s/   ø€ ô ‘Ñ" 5Ó)ˆØ˜w™×-Ñ-Ó/ˆˆg‰Øˆr   c                ó  •— t         ‰| �  |«       | j                  j                  rt	        d«       | j                  j                  s| j                  j
                  rt        j                  | _	        yt        j                  | _	        y)zŸInitialize metrics and select mask processing function based on save_json flag.

        Args:
            model (torch.nn.Module): Model to validate.
        zfaster-coco-eval>=1.6.7N)
r   Úinit_metricsr   Ú	save_jsonr	   Úsave_txtr   Úprocess_mask_nativeÚprocess_maskr   )r   Úmodelr   s     €r   r#   z"SegmentationValidator.init_metricsA   sZ   ø€ ô 	‰Ñ˜UÔ#Ø�9‰9×ÒÜÐ8Ô9à26·)±)×2EÒ2EÈÏÉ×I[ÒI[”s×.Ñ.ˆ�Ôad×aqÑaqˆ�r   c                ó   — ddz  S )z5Return a formatted description of evaluation metrics.z,%22s%11s%11s%11s%11s%11s%11s%11s%11s%11s%11s)ÚClassÚImagesÚ	InstanceszBox(PÚRÚmAP50ú	mAP50-95)zMask(Pr-   r.   r/   © )r   s    r   Úget_desczSegmentationValidator.get_descM   s   € à$ð )
ñ 
ð 	
r   c                ó  •— t        |d   t        «      r|d   d   n|d   }t        ‰| �  |d   «      }|j                  dd D �cg c]  }d|z  ‘Œ	 }}t        |«      D ]¥  \  }}|j                  d«      }|j                  d   r| j                  ||   ||d   |¬«      nat        j                  dg| j                  t        j                  u r|n|j                  dd ¢­t        j                  |d   j                  ¬	«      |d
<   Œ§ |S c c}w )a  Post-process YOLO predictions and return output detections with proto.

        Args:
            preds (list[torch.Tensor]): Raw predictions from the model.

        Returns:
            (list[dict[str, torch.Tensor]]): Processed detection predictions with masks.
        r   é   é   Né   ÚextraÚbboxes)Úshape)ÚdtypeÚdevicer   )Ú
isinstanceÚtupler   Úpostprocessr8   Ú	enumerateÚpopr   ÚtorchÚzerosr   r&   Úuint8r:   )	r   ÚpredsÚprotoÚxÚimgszÚiÚpredÚcoefficientr   s	           €r   r=   z!SegmentationValidator.postprocess]   s  ø€ ô  *¨%°©(´EÔ:��a‘˜’ÀÀaÁˆÜ‘Ñ# E¨!¡HÓ-ˆØ %§¡¨A¨B Ö0˜1��Q“Ð0ˆÐ0Ü  Ó'ò 
	‰GˆAˆtØŸ(™( 7Ó+ˆKð ×$Ñ$ QÒ'ð —‘˜U 1™X {°D¸±NÈ%�ÔPä—[‘[ØÐa 4§<¡<´3×3JÑ3JÑ#J™%ÐPU×P[ÑP[Ð\]Ð\^ÐP_ÑaÜŸ+™+Ø ™>×0Ñ0ôð �ŠMð
	ð ˆùò 1s   ÁDc                ó:  •— t         ‰	| �  ||«      }|d   j                  d   }| j                  j                  rR|d   |   }t        j                  d|dz   |j                  ¬«      j                  |dd«      }||k(  j                  «       }n|d   |d   |k(     }|ru|d   D �cg c]%  }| j                  t        j                  u r|n|dz  ‘Œ' }}|j                  dd	 |k7  r0t        j                  |d	   |d
d¬«      d   }|j                  d«      }||d<   |S c c}w )a:  Prepare a batch for validation by processing images and targets.

        Args:
            si (int): Sample index within the batch.
            batch (dict[str, Any]): Batch data containing images and annotations.

        Returns:
            (dict[str, Any]): Prepared batch with processed annotations.
        Úclsr   r   r3   )r:   Ú	batch_idxrF   r5   NÚbilinearF)ÚmodeÚalign_cornersg      à?)r   Ú_prepare_batchr8   r   Úoverlap_maskr@   Úaranger:   Úviewr    r   r   r&   ÚFÚinterpolateÚgt_)
r   Úsir!   Úprepared_batchÚnlr   ÚindexÚsÚ	mask_sizer   s
            €r   rP   z$SegmentationValidator._prepare_batchv   s%  ø€ ô ™Ñ/°°EÓ:ˆØ˜EÑ"×(Ñ(¨Ñ+ˆØ�9‰9×!Ò!Ø˜'‘N 2Ñ&ˆEÜ—L‘L  B¨¡F°5·<±<Ô@×EÑEÀbÈ!ÈQÓOˆEØ˜e‘^×*Ñ*Ó,‰Eà˜'‘N 5¨Ñ#5¸Ñ#;Ñ<ˆEÙØ[iÐjqÑ[rÖsÐVW˜dŸl™l¬c×.EÑ.EÑE™È1ÐPQÉ6ÑQÐsˆIÐsØ�{‰{˜1˜2ˆ )Ò+ÜŸ™ e¨D¡k°9À:Ð]bÔcÐdeÑf�ØŸ	™	 #›�Ø"'ˆ�wÑØÐùò ts   Â#*Dc                ól   •— t         ‰| �  «        | j                  | j                  j                  «       y)zGather stats from all GPUs.N)r   Úgather_statsÚ_gather_image_metricsr   Úseg)r   r   s    €r   r^   z"SegmentationValidator.gather_stats�   s&   ø€ ä‰ÑÔØ×"Ñ" 4§<¡<×#3Ñ#3Õ4r   c                óö  •— t         ‰| �  ||«      }|d   }|j                  d   dk(  s|d   j                  d   dk(  r8t        j                  |d   j                  d   | j
                  ft        ¬«      }npt        |d   j                  d«      |d   j                  d«      j                  «       «      }| j                  |d   ||«      j                  «       j                  «       }|j                  d|i«       |S )aÎ  Compute correct prediction matrix for a batch based on bounding boxes and optional masks.

        Args:
            preds (dict[str, torch.Tensor]): Dictionary containing predictions with keys like 'cls' and 'masks'.
            batch (dict[str, Any]): Dictionary containing batch data with keys like 'cls' and 'masks'.

        Returns:
            (dict[str, np.ndarray]): A dictionary containing correct prediction matrices including 'tp_m' for mask IoU.

        Examples:
            >>> preds = {"cls": torch.tensor([1, 0]), "masks": torch.rand(2, 640, 640), "bboxes": torch.rand(2, 4)}
            >>> batch = {"cls": torch.tensor([1, 0]), "masks": torch.rand(2, 640, 640), "bboxes": torch.rand(2, 4)}
            >>> correct_preds = validator._process_batch(preds, batch)

        Notes:
            - This method computes IoU between predicted and ground truth masks.
            - Overlapping masks are handled based on the overlap_mask argument setting.
        rK   r   ©r9   r   r3   Útp_m)r   Ú_process_batchr8   ÚnprA   ÚniouÚboolr   Úflattenr    Úmatch_predictionsÚcpuÚnumpyÚupdate)r   rC   r!   ÚtpÚgt_clsrc   Úiour   s          €r   rd   z$SegmentationValidator._process_batch•   sß   ø€ ô& ‰WÑ# E¨5Ó1ˆØ�u‘ˆØ�<‰<˜‰?˜aÒ 5¨¡<×#5Ñ#5°aÑ#8¸AÒ#=Ü—8‘8˜U 5™\×/Ñ/°Ñ2°D·I±IÐ>ÄdÔK‰Dä˜5 ™>×1Ñ1°!Ó4°e¸G±n×6LÑ6LÈQÓ6O×6UÑ6UÓ6WÓXˆCØ×)Ñ)¨%°©,¸ÀÓD×HÑHÓJ×PÑPÓRˆDØ
�	‰	�6˜4�.Ô!Øˆ	r   c                ó¬  •— |D ]§  }|d   }|j                   d   | j                  j                  kD  r-t        j                  d| j                  j                  › d�«       t        j                  |d| j                  j                   t
        j                  ¬«      j                  «       |d<   Œ© t        ‰| �)  |||| j                  j                  ¬«       y)a  Plot batch predictions with masks and bounding boxes.

        Args:
            batch (dict[str, Any]): Batch containing images and annotations.
            preds (list[dict[str, torch.Tensor]]): List of predictions from the model.
            ni (int): Batch index.
        r   r   z&Limiting validation plots to 'max_det=z' items.Nrb   )Úmax_det)r8   r   rq   r   Úwarningr@   Ú	as_tensorrB   rj   r   Úplot_predictions)r   r!   rC   ÚniÚpr   r   s         €r   rt   z&SegmentationValidator.plot_predictions²   s°   ø€ ð ò 	^ˆAØ�g‘JˆEØ�{‰{˜1‰~ §	¡	× 1Ñ 1Ò1Ü—‘Ð!GÈÏ	É	×HYÑHYÐGZÐZbÐcÔdÜŸ™¨Ð/B°·±×1BÑ1BÐ)CÌ5Ï;É;ÔW×[Ñ[Ó]ˆAˆgŠJð		^ô
 	‰Ñ  ¨¨r¸4¿9¹9×;LÑ;LÐ ÕMr   c                ó€  — ddl m}  |t        j                  |d   |d   ft        j                  ¬«      d| j
                  t        j                  |d   |d   j                  d«      |d	   j                  d«      gd¬
«      t        j                  |d   t        j                  ¬«      ¬«      j                  ||¬«       y)a¡  Save YOLO detections to a txt file in normalized coordinates in a specific format.

        Args:
            predn (dict[str, torch.Tensor]): Prediction dictionary containing 'bboxes', 'conf', 'cls', and 'masks' keys.
            save_conf (bool): Whether to save confidence scores.
            shape (tuple[int, int]): Shape of the original image.
            file (Path): File path to save the detections.
        r   )ÚResultsr3   rb   Nr7   ÚconféÿÿÿÿrK   )Údimr   )ÚpathÚnamesÚboxesr   )Ú	save_conf)Úultralytics.engine.resultsrx   re   rA   rB   r}   r@   ÚcatÚ	unsqueezers   r%   )r   Úprednr   r8   Úfilerx   s         r   Úsave_one_txtz"SegmentationValidator.save_one_txtÁ   s˜   € õ 	7áÜ�H‰H�e˜A‘h  a¡Ð)´·±Ô:ØØ—*‘*Ü—)‘)˜U 8™_¨e°F©m×.EÑ.EÀbÓ.IÈ5ÐQVÉ<×KaÑKaÐbdÓKeÐfÐlmÔnÜ—/‘/ %¨¡.¼¿¹ÔDô	
÷ ‰(�4 9ˆ(Õ
-r   c                óœ  •— dd„}dd„}|d   j                  dd«      j                  «       j                  t        |d   «      d«      }|d   j                  dd \  }} ||«      }g }	|D ]  }
|	j                  ||g ||
«      dœ«       Œ  t        ‰| �  ||«       t        |	«      D ]$  \  }}|| j                  t        |	«       |z      d	<   Œ& y
)a'  Save one JSON result for COCO evaluation.

        Args:
            predn (dict[str, torch.Tensor]): Predictions containing bboxes, masks, confidence scores, and classes.
            pbatch (dict[str, Any]): Batch dictionary containing 'imgsz', 'ori_shape', 'ratio_pad', and 'im_file'.
        c                ó.  — g }t        t        | «      «      D ]l  }t        | |   «      }|dkD  r|t        | |dz
     «      z  }	 |dz  }|dz  }|dz  r|dk7  n|dk7  }|r|dz  }|dz  }|j                  t	        |«      «       |sŒmŒC d	j                  |«      S )
zæConverts the RLE object into a compact string representation. Each count is delta-encoded and
            variable-length encoded as a string.

            Args:
                counts (list[int]): List of RLE counts.
            r4   é   é   é   rz   r   é    é0   Ú )ÚrangeÚlenÚintÚappendÚchrÚjoin)ÚcountsÚresultrG   rE   ÚcÚmores         r   Ú	to_stringz5SegmentationValidator.pred_to_json.<locals>.to_stringÜ   s¼   € ð ˆFäœ3˜v›;Ó'ò �Ü˜˜q™	“N�ð �q’5Øœ˜V A¨¡E™]Ó+Ñ+�Að Ø˜D™�AØ˜!‘G�Að *+¨Tª˜A šG¸¸a¹�DÙØ˜T™	˜Ø˜‘G�AØ—M‘M¤# a£&Ô)ÙØð ðð, —7‘7˜6“?Ð"r   c                ó>  — | dd…dd…f   | dd…dd…f   k7  }t        j                  |«      \  }}|dz   }g }t        | j                  d   «      D ]Ë  }|||k(     }t	        |«      rxt        j
                  |«      j                  «       }|j                  d|d   j                  «       «       |j                  t	        | |   «      |d   j                  «       z
  «       nt	        | |   «      g}| |   d   j                  «       dk(  rdg|¢}|j                  |«       ŒÍ |S )aI  Convert multiple binary masks using Run-Length Encoding (RLE).

            Args:
                pixels (torch.Tensor): A 2D tensor where each row represents a flattened binary mask with shape [N,
                    H*W].

            Returns:
                (list[list[int]]): A list of RLE counts for each mask.
            Nr3   rz   r   )
r@   ÚwhererŽ   r8   r�   ÚdiffÚtolistÚinsertÚitemr‘   )ÚpixelsÚtransitionsÚrow_idxÚcol_idxr”   rG   Ú	positionsÚcounts           r   Úmulti_encodez8SegmentationValidator.pred_to_json.<locals>.multi_encodeý   s  € ð !¢ A¡B ™-¨6²!°S°b°S°&©>Ñ9ˆKÜ$Ÿ{™{¨;Ó7ÑˆG�WØ ‘kˆGð ˆFÜ˜6Ÿ<™<¨™?Ó+ò %�Ø# G¨q¡LÑ1�	Ü�y”>Ü!ŸJ™J yÓ1×8Ñ8Ó:�EØ—L‘L  I¨a¡L×$5Ñ$5Ó$7Ô8Ø—L‘L¤ V¨A¡Y£°)¸B±-×2DÑ2DÓ2FÑ!FÕGä  ¨¡›^Ð,�Eð ˜!‘9˜Q‘<×$Ñ$Ó&¨!Ò+Ø˜K ˜K�EØ—‘˜eÕ$ð%ð ˆMr   r   r4   r3   rz   é   )Úsizer”   ÚsegmentationN)r”   ú	list[int]ÚreturnÚstr)rŸ   ztorch.Tensorrª   r©   )
Ú	transposeÚ
contiguousrS   r�   r8   r‘   r   Úpred_to_jsonr>   Újdict)r   rƒ   Úpbatchr˜   r¥   Ú
pred_masksÚhÚwr”   Úrlesr–   rG   Úrr   s                €r   r®   z"SegmentationValidator.pred_to_jsonÔ   sÞ   ø€ ó	#óB	ð@ ˜7‘^×-Ñ-¨a°Ó3×>Ñ>Ó@×EÑEÄcÈ%ÐPWÉ.ÓFYÐ[]Ó^ˆ
Ø�W‰~×#Ñ# A aÐ(‰ˆˆ1Ù˜jÓ)ˆØˆØò 	BˆAØ�K‰K ! Q ±9¸Q³<Ñ@ÕAð	Bä‰Ñ˜U FÔ+Ü˜d“Oò 	;‰DˆAˆqØ9:ˆD�J‰Jœ˜D›	�z A‘~Ñ& ~Ò6ñ	;r   c                ó–   •— i t         ‰| �  ||«      ¥dt        j                  |d   d   |d   |d   ¬«      d   j	                  «       i¥S )z.Scales predictions to the original image size.r   NÚ	ori_shapeÚ	ratio_pad)r¸   r   )r   Úscale_predsr   Úscale_masksÚbyte)r   rƒ   r°   r   s      €r   r¹   z!SegmentationValidator.scale_preds'  s\   ø€ ð
Ü‰gÑ! %¨Ó0ð
à”S—_‘_ U¨7¡^°DÑ%9¸6À+Ñ;NÐZ`ÐalÑZmÔnØñç‰d‹fñ	
ð 	
r   c                óÈ   •— | j                   dz  }| j                  d   dz  | j                  rdnd| j                  j                  › d�z  }t
        ‰| �  |||ddgd	d
g¬«      S )z;Return COCO-style instance segmentation evaluation metrics.zpredictions.jsonr|   r   zinstances_val2017.jsonÚlvis_v1_z.jsonÚbboxÚsegmÚBoxÚMask)Úsuffix)r   ÚdataÚis_cocor   Úsplitr   Úcoco_evaluate)r   ÚstatsÚ	pred_jsonÚ	anno_jsonr   s       €r   Ú	eval_jsonzSegmentationValidator.eval_json0  sy   ø€ à—M‘MÐ$6Ñ6ˆ	à�I‰I�fÑØñà+/¯<ª<Ñ'¸xÈÏ	É	ÏÉÐGXÐX]Ð=^ñ`ð 	ô
 ‰wÑ$ U¨I°yÀ6È6ÐBRÐ\aÐciÐ[jÐ$ÓkÐkr   )NNNN)r   zdict | Nonerª   ÚNone)r!   údict[str, Any]rª   rÌ   )r(   ztorch.nn.Modulerª   rË   )rª   r«   )rC   zlist[torch.Tensor]rª   úlist[dict[str, torch.Tensor]])rW   r�   r!   rÌ   rª   rÌ   )rª   rË   )rC   údict[str, torch.Tensor]r!   rÌ   rª   zdict[str, np.ndarray])r!   rÌ   rC   rÍ   ru   r�   rª   rË   )
rƒ   rÎ   r   rg   r8   ztuple[int, int]r„   r   rª   rË   )rƒ   rÎ   r°   rÌ   rª   rË   )rƒ   rÎ   r°   rÌ   rª   rÎ   )rÇ   rÌ   rª   rÌ   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r#   r1   r=   rP   r^   rd   rt   r…   r®   r¹   rÊ   Ú__classcell__)r   s   @r   r   r      sT   ø„ ñö&(õõ
ró
õ õ2õ45õ
õ:Nó.õ&Q;õf
÷lñ lr   r   )Ú
__future__r   Úpathlibr   Útypingr   rk   re   r@   Útorch.nn.functionalÚnnÚ
functionalrT   Úultralytics.models.yolo.detectr   Úultralytics.utilsr   r   Úultralytics.utils.checksr	   Úultralytics.utils.metricsr
   r   r   r0   r   r   ú<module>rÞ      s9   ðõ #å Ý ã Û ß Ð å =ß )Ý 7ß >ôflÐ.õ flr   