Ë
    Fêñi)  ã                  óš   — d dl mZ d dlmZ d dlmZ d dl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 d dlmZmZ d d	lmZ  G d
„ de«      Zy)é    )Úannotations)ÚPath)ÚAnyN)ÚClassificationDatasetÚbuild_dataloader)ÚBaseValidator)ÚLOGGERÚRANK)ÚClassifyMetricsÚConfusionMatrix)Úplot_imagesc                  óŽ   ‡ — e Zd ZdZddˆ fd„Zdd„Zdd„Zdd„Zdd„Zdd„Z	dd„Z
dd	„Zdd
„Zdd„Zdd„Zdd„Zdd„Zdd„Zˆ xZS )ÚClassificationValidatoraì  A class extending the BaseValidator class for validation based on a classification model.

    This validator handles the validation process for classification models, including metrics calculation, confusion
    matrix generation, and visualization of results.

    Attributes:
        targets (list[torch.Tensor]): Ground truth class labels.
        pred (list[torch.Tensor]): Model predictions.
        metrics (ClassifyMetrics): Object to calculate and store classification metrics.
        names (dict): Mapping of class indices to class names.
        nc (int): Number of classes.
        confusion_matrix (ConfusionMatrix): Matrix to evaluate model performance across classes.

    Methods:
        get_desc: Return a formatted string summarizing classification metrics.
        init_metrics: Initialize confusion matrix, class names, and tracking containers.
        preprocess: Preprocess input batch by moving data to device.
        update_metrics: Update running metrics with model predictions and batch targets.
        finalize_metrics: Finalize metrics including confusion matrix and processing speed.
        postprocess: Extract the primary prediction from model output.
        get_stats: Calculate and return a dictionary of metrics.
        build_dataset: Create a ClassificationDataset instance for validation.
        get_dataloader: Build and return a data loader for classification validation.
        print_results: Print evaluation metrics for the classification model.
        plot_val_samples: Plot validation image samples with their ground truth labels.
        plot_predictions: Plot images with their predicted class labels.

    Examples:
        >>> from ultralytics.models.yolo.classify import ClassificationValidator
        >>> args = dict(model="yolo26n-cls.pt", data="imagenet10")
        >>> validator = ClassificationValidator(args=args)
        >>> validator()

    Notes:
        Torchvision classification models can also be passed to the 'model' argument, i.e. model='resnet18'.
    c                ó†   •— t         ‰| �  ||||«       d| _        d| _        d| j                  _        t        «       | _        y)aá  Initialize ClassificationValidator with dataloader, save directory, and other parameters.

        Args:
            dataloader (torch.utils.data.DataLoader, optional): DataLoader to use for validation.
            save_dir (str | Path, optional): Directory to save results.
            args (dict, optional): Arguments containing model and validation configuration.
            _callbacks (dict, optional): Dictionary of callback functions to be called during validation.
        NÚclassify)ÚsuperÚ__init__ÚtargetsÚpredÚargsÚtaskr   Úmetrics)ÚselfÚ
dataloaderÚsave_dirr   Ú
_callbacksÚ	__class__s        €úf/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/yolo/classify/val.pyr   z ClassificationValidator.__init__8   s;   ø€ ô 	‰Ñ˜ X¨t°ZÔ@ØˆŒØˆŒ	Ø#ˆ�	‰	ŒÜ&Ó(ˆ�ó    c                ó   — ddz  S )z=Return a formatted string summarizing classification metrics.z%22s%11s%11s)ÚclassesÚtop1_accÚtop5_acc© ©r   s    r   Úget_descz ClassificationValidator.get_descG   s   € à#Ð'JÑJÐJr   c                ó¬   — |j                   | _         t        |j                   «      | _        g | _        g | _        t        |j                   ¬«      | _        y)z^Initialize confusion matrix, class names, and tracking containers for predictions and targets.)ÚnamesN)r(   ÚlenÚncr   r   r   Úconfusion_matrix)r   Úmodels     r   Úinit_metricsz$ClassificationValidator.init_metricsK   s<   € à—[‘[ˆŒ
Ü�e—k‘kÓ"ˆŒØˆŒ	ØˆŒÜ /°e·k±kÔ BˆÕr   c                ól  — |d   j                  | j                  | j                  j                  dk(  ¬«      |d<   | j                  j                  r|d   j	                  «       n|d   j                  «       |d<   |d   j                  | j                  | j                  j                  dk(  ¬«      |d<   |S )zTPreprocess input batch by moving data to device and converting to appropriate dtype.ÚimgÚcuda)Únon_blockingÚcls)ÚtoÚdeviceÚtyper   ÚhalfÚfloat)r   Úbatchs     r   Ú
preprocessz"ClassificationValidator.preprocessS   s’   € à˜U‘|—‘ t§{¡{ÀÇÁ×AQÑAQÐU[ÑA[�Ó\ˆˆe‰Ø.2¯i©i¯nªn�u˜U‘|×(Ñ(Ô*À%ÈÁ,×BTÑBTÓBVˆˆe‰Ø˜U‘|—‘ t§{¡{ÀÇÁ×AQÑAQÐU[ÑA[�Ó\ˆˆe‰Øˆr   c                ó”  — t        t        | j                  «      d«      }| j                  j	                  |j                  dd¬«      dd…d|…f   j                  t        j                  «      j                  «       «       | j                  j	                  |d   j                  t        j                  «      j                  «       «       y)aî  Update running metrics with model predictions and batch targets.

        Args:
            preds (torch.Tensor): Model predictions, typically logits or probabilities for each class.
            batch (dict): Batch data containing images and class labels.

        Notes:
            This method appends the top-N predictions (sorted by confidence in descending order) to the
            prediction list for later evaluation. N is limited to the minimum of 5 and the number of classes.
        é   é   T)Ú
descendingNr2   )Úminr)   r(   r   ÚappendÚargsortr5   ÚtorchÚint32Úcpur   )r   Úpredsr8   Ún5s       r   Úupdate_metricsz&ClassificationValidator.update_metricsZ   sŠ   € ô ”�T—Z‘Z“ !Ó$ˆØ�	‰	×Ñ˜Ÿ™ q°T˜Ó:º1¸c¸r¸c¸6ÑB×GÑGÌÏÉÓT×XÑXÓZÔ[Ø�‰×Ñ˜E %™L×-Ñ-¬e¯k©kÓ:×>Ñ>Ó@ÕAr   c                ó¤  — | j                   j                  | j                  | j                  «       | j                  j
                  r9dD ]4  }| j                   j                  | j                  || j                  ¬«       Œ6 | j                  | j                  _	        | j                  | j                  _        | j                   | j                  _         y)aœ  Finalize metrics including confusion matrix and processing speed.

        Examples:
            >>> validator = ClassificationValidator()
            >>> validator.pred = [torch.tensor([[0, 1, 2]])]  # Top-3 predictions for one sample
            >>> validator.targets = [torch.tensor([0])]  # Ground truth class
            >>> validator.finalize_metrics()
            >>> print(validator.metrics.confusion_matrix)  # Access the confusion matrix

        Notes:
            This method processes the accumulated predictions and targets to generate the confusion matrix,
            optionally plots it, and updates the metrics object with speed information.
        )TF)r   Ú	normalizeÚon_plotN)r+   Úprocess_cls_predsr   r   r   ÚplotsÚplotr   rI   Úspeedr   )r   rH   s     r   Úfinalize_metricsz(ClassificationValidator.finalize_metricsi   s–   € ð 	×Ñ×/Ñ/°·	±	¸4¿<¹<ÔHØ�9‰9�?Š?Ø(ò n�	Ø×%Ñ%×*Ñ*°D·M±MÈYÐ`d×`lÑ`lÐ*Õmðnà!ŸZ™Zˆ�‰ÔØ $§¡ˆ�‰ÔØ(,×(=Ñ(=ˆ�‰Õ%r   c                ó<   — t        |t        t        f«      r|d   S |S )zSExtract the primary prediction from model output if it's in a list or tuple format.r   )Ú
isinstanceÚlistÚtuple)r   rD   s     r   Úpostprocessz#ClassificationValidator.postprocess   s   € ä% e¬d´E¨]Ô;ˆu�Q‰xÐFÀÐFr   c                óŽ   — | j                   j                  | j                  | j                  «       | j                   j                  S )zSCalculate and return a dictionary of metrics by processing targets and predictions.)r   Úprocessr   r   Úresults_dictr%   s    r   Ú	get_statsz!ClassificationValidator.get_statsƒ   s.   € à�‰×Ñ˜TŸ\™\¨4¯9©9Ô5Ø�|‰|×(Ñ(Ð(r   c                ó,  — t         dk(  r±dgt        j                  «       z  }dgt        j                  «       z  }t        j                  | j                  |d¬«       t        j                  | j
                  |d¬«       |D ��cg c]  }|D ]  }|‘Œ Œ c}}| _        |D ��cg c]  }|D ]  }|‘Œ Œ c}}| _        yt         dkD  rEt        j                  | j                  dd¬«       t        j                  | j
                  dd¬«       yyc c}}w c c}}w )zGather stats from all GPUs.r   N)Údst)r
   ÚdistÚget_world_sizeÚgather_objectr   r   )r   Úgathered_predsÚgathered_targetsÚrankr   r   s         r   Úgather_statsz$ClassificationValidator.gather_statsˆ   sß   € ä�1Š9Ø"˜V¤d×&9Ñ&9Ó&;Ñ;ˆNØ $˜v¬×(;Ñ(;Ó(=Ñ=ÐÜ×Ñ˜tŸy™y¨.¸aÕ@Ü×Ñ˜tŸ|™|Ð-=À1ÕEØ*8×J $ÀTÒJ¸TšÐJ˜ÓJˆDŒIØ0@×U¨ÐPTÒUÀWšGÐU˜GÓUˆD�LÜ�AŠXÜ×Ñ˜tŸy™y¨$°AÕ6Ü×Ñ˜tŸ|™|¨T°qÖ9ð ùó KùÛUs   ÂD
Â!Dc                ó\   — t        || j                  d| j                  j                  ¬«      S )z7Create a ClassificationDataset instance for validation.F)Úrootr   ÚaugmentÚprefix)r   r   Úsplit)r   Úimg_paths     r   Úbuild_datasetz%ClassificationValidator.build_dataset•   s$   € ä$¨(¸¿¹ÈEÐZ^×ZcÑZc×ZiÑZiÔjÐjr   c                ój   — | j                  |«      }t        ||| j                  j                  d¬«      S )aP  Build and return a data loader for classification validation.

        Args:
            dataset_path (str | Path): Path to the dataset directory.
            batch_size (int): Number of samples per batch.

        Returns:
            (torch.utils.data.DataLoader): DataLoader object for the classification validation dataset.
        éÿÿÿÿ)r_   )rg   r   r   Úworkers)r   Údataset_pathÚ
batch_sizeÚdatasets       r   Úget_dataloaderz&ClassificationValidator.get_dataloader™   s/   € ð ×$Ñ$ \Ó2ˆÜ ¨°T·Y±Y×5FÑ5FÈRÔPÐPr   c                óÔ   — ddt        | j                  j                  «      z  z   }t        j                  |d| j                  j
                  | j                  j                  fz  «       y)z6Print evaluation metrics for the classification model.z%22sz%11.3gÚallN)r)   r   Úkeysr	   ÚinfoÚtop1Útop5)r   Úpfs     r   Úprint_resultsz%ClassificationValidator.print_results¦   sL   € à�h¤ T§\¡\×%6Ñ%6Ó!7Ñ7Ñ7ˆÜ�‰�B˜% §¡×!2Ñ!2°D·L±L×4EÑ4EÐFÑFÕGr   c                ó¼   — t        j                  |d   j                  d   «      |d<   t        || j                  d|› d�z  | j
                  | j                  ¬«       y)aê  Plot validation image samples with their ground truth labels.

        Args:
            batch (dict[str, Any]): Dictionary containing batch data with 'img' (images) and 'cls' (class labels).
            ni (int): Batch index used for naming the output file.

        Examples:
            >>> validator = ClassificationValidator()
            >>> batch = {"img": torch.rand(16, 3, 224, 224), "cls": torch.randint(0, 10, (16,))}
            >>> validator.plot_val_samples(batch, 0)
        r/   r   Ú	batch_idxÚ	val_batchz_labels.jpg)ÚlabelsÚfnamer(   rI   N)rA   ÚarangeÚshaper   r   r(   rI   )r   r8   Únis      r   Úplot_val_samplesz(ClassificationValidator.plot_val_samples«   sT   € ô #Ÿ\™\¨%°©,×*<Ñ*<¸QÑ*?Ó@ˆˆkÑÜØØ—-‘- I¨b¨T°Ð"=Ñ=Ø—*‘*Ø—L‘Lö		
r   c           	     ó*  — t        |d   t        j                  |d   j                  d   «      t        j                  |d¬«      t        j
                  |d¬«      ¬«      }t        || j                  d|› d�z  | j                  | j                  ¬«       y	)
a\  Plot images with their predicted class labels and save the visualization.

        Args:
            batch (dict[str, Any]): Batch data containing images and other information.
            preds (torch.Tensor): Model predictions with shape (batch_size, num_classes).
            ni (int): Batch index used for naming the output file.

        Examples:
            >>> validator = ClassificationValidator()
            >>> batch = {"img": torch.rand(16, 3, 224, 224)}
            >>> preds = torch.rand(16, 10)  # 16 images, 10 classes
            >>> validator.plot_predictions(batch, preds, 0)
        r/   r   r<   )Údim)r/   rx   r2   Úconfry   z	_pred.jpg)r{   r(   rI   N)
ÚdictrA   r|   r}   ÚargmaxÚamaxr   r   r(   rI   )r   r8   rD   r~   Úbatched_predss        r   Úplot_predictionsz(ClassificationValidator.plot_predictions¿   s|   € ô Ø�e‘Ü—l‘l 5¨¡<×#5Ñ#5°aÑ#8Ó9Ü—‘˜U¨Ô*Ü—‘˜E qÔ)ô	
ˆô 	ØØ—-‘- I¨b¨T°Ð";Ñ;Ø—*‘*Ø—L‘Lö		
r   )NNNN)r   zdict | NoneÚreturnÚNone)rˆ   Ústr)r,   ztorch.nn.Modulerˆ   r‰   )r8   údict[str, Any]rˆ   r‹   )rD   útorch.Tensorr8   r‹   rˆ   r‰   )rˆ   r‰   )rD   z7torch.Tensor | list[torch.Tensor] | tuple[torch.Tensor]rˆ   rŒ   )rˆ   zdict[str, float])rf   rŠ   rˆ   r   )rk   z
Path | strrl   Úintrˆ   ztorch.utils.data.DataLoader)r8   r‹   r~   r�   rˆ   r‰   )r8   r‹   rD   rŒ   r~   r�   rˆ   r‰   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r&   r-   r9   rF   rN   rS   rW   r`   rg   rn   rv   r   r‡   Ú__classcell__)r   s   @r   r   r      sV   ø„ ñ#öJ)óKóCóóBó>ó,Gó)ó
:ókóQóHó

÷(
r   r   )Ú
__future__r   Úpathlibr   Útypingr   rA   Útorch.distributedÚdistributedrZ   Úultralytics.datar   r   Úultralytics.engine.validatorr   Úultralytics.utilsr	   r
   Úultralytics.utils.metricsr   r   Úultralytics.utils.plottingr   r   r$   r   r   ú<module>r�      s3   ðõ #å Ý ã Ý  ç DÝ 6ß *ß FÝ 2ôF
˜mõ F
r   