Ë
    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mZ d dl	m
Z
 d dlmZ d dlmZ d d	lmZmZmZ d d
lmZ d dlmZmZ  G d„ de
«      Zy)é    )Úannotations)Úcopy)ÚAnyN)ÚClassificationDatasetÚbuild_dataloader)ÚBaseTrainer)Úyolo)ÚClassificationModel)ÚDEFAULT_CFGÚLOGGERÚRANK)Úplot_images)Úis_parallelÚtorch_distributed_zero_firstc                  ó‚   ‡ — e Zd ZdZeddfdˆ fd„Zd„ Zddd„Zˆ fd„Zddd„Z	ddd„Z
dd	„Zdd
„Zd„ Zddd„Zdd„Zˆ xZS )ÚClassificationTrainera  A trainer class extending BaseTrainer for training image classification models.

    This trainer handles the training process for image classification tasks, supporting both YOLO classification models
    and torchvision models with comprehensive dataset handling and validation.

    Attributes:
        model (ClassificationModel): The classification model to be trained.
        data (dict[str, Any]): Dictionary containing dataset information including class names and number of classes.
        loss_names (list[str]): Names of the loss functions used during training.
        validator (ClassificationValidator): Validator instance for model evaluation.

    Methods:
        set_model_attributes: Set the model's class names from the loaded dataset.
        get_model: Return a modified PyTorch model configured for training.
        setup_model: Load, create or download model for classification.
        build_dataset: Create a ClassificationDataset instance.
        get_dataloader: Return PyTorch DataLoader with transforms for image preprocessing.
        preprocess_batch: Preprocess a batch of images and classes.
        progress_string: Return a formatted string showing training progress.
        get_validator: Return an instance of ClassificationValidator.
        label_loss_items: Return a loss dict with labeled training loss items.
        final_eval: Evaluate trained model and save validation results.
        plot_training_samples: Plot training samples with their annotations.

    Examples:
        Initialize and train a classification model
        >>> from ultralytics.models.yolo.classify import ClassificationTrainer
        >>> args = dict(model="yolo26n-cls.pt", data="imagenet10", epochs=3)
        >>> trainer = ClassificationTrainer(overrides=args)
        >>> trainer.train()
    Nc                óf   •— |€i }d|d<   |j                  d«      €d|d<   t        ‰| �	  |||«       y)aŒ  Initialize a ClassificationTrainer object.

        Args:
            cfg (dict[str, Any], optional): Default configuration dictionary containing training parameters.
            overrides (dict[str, Any], optional): Dictionary of parameter overrides for the default configuration.
            _callbacks (dict, optional): Dictionary of callback functions to be executed during training.
        NÚclassifyÚtaskÚimgszéà   )ÚgetÚsuperÚ__init__)ÚselfÚcfgÚ	overridesÚ
_callbacksÚ	__class__s       €úh/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/yolo/classify/train.pyr   zClassificationTrainer.__init__4   sD   ø€ ð ÐØˆIØ&ˆ	�&ÑØ�=‰=˜Ó!Ð)Ø!$ˆI�gÑÜ‰Ñ˜˜i¨Õ4ó    c                ó@   — | j                   d   | j                  _        y)z9Set the YOLO model's class names from the loaded dataset.ÚnamesN)ÚdataÚmodelr#   ©r   s    r    Úset_model_attributesz*ClassificationTrainer.set_model_attributesC   s   € àŸ9™9 WÑ-ˆ�
‰
Õr!   c                ó  — t        || j                  d   | j                  d   |xr	 t        dk(  ¬«      }|r|j                  |«       |j	                  «       D ]‹  }| j
                  j                  st        |d«      r|j                  «        t        |t        j                  j                  «      sŒZ| j
                  j                  sŒq| j
                  j                  |_        Œ� |j                  «       D ]	  }d|_        Œ |S )aˆ  Return a modified PyTorch model configured for training YOLO classification.

        Args:
            cfg (Any, optional): Model configuration.
            weights (Any, optional): Pre-trained model weights.
            verbose (bool, optional): Whether to display model information.

        Returns:
            (ClassificationModel): Configured PyTorch model for classification.
        ÚncÚchannelséÿÿÿÿ)r)   ÚchÚverboseÚreset_parametersT)r
   r$   r   ÚloadÚmodulesÚargsÚ
pretrainedÚhasattrr.   Ú
isinstanceÚtorchÚnnÚDropoutÚdropoutÚpÚ
parametersÚrequires_grad)r   r   Úweightsr-   r%   Úmr9   s          r    Ú	get_modelzClassificationTrainer.get_modelG   sÐ   € ô $ C¨D¯I©I°d©OÀÇ	Á	È*Ñ@UÐ_fÒ_uÔkoÐsuÑkuÔvˆÙØ�J‰J�wÔà—‘“ò 	(ˆAØ—9‘9×'Ò'¬G°AÐ7IÔ,JØ×"Ñ"Ô$Ü˜!œUŸX™X×-Ñ-Õ.°4·9±9×3DÓ3DØ—i‘i×'Ñ'�•ð		(ð
 ×!Ñ!Ó#ò 	#ˆAØ"ˆA�Oð	#àˆr!   c                óp  •— ddl }t        | j                  «      |j                  j                  v rJ |j                  j                  | j                     | j
                  j                  rdnd¬«      | _        d}nt        ‰| �!  «       }t        j                  | j                  | j                  d   «       |S )z–Load, create or download model for classification tasks.

        Returns:
            (Any): Model checkpoint if applicable, otherwise None.
        r   NÚIMAGENET1K_V1)r<   r)   )ÚtorchvisionÚstrr%   ÚmodelsÚ__dict__r1   r2   r   Úsetup_modelr
   Úreshape_outputsr$   )r   rA   Úckptr   s      €r    rE   z!ClassificationTrainer.setup_model_   s�   ø€ ó 	äˆt�z‰z‹?˜k×0Ñ0×9Ñ9Ñ9Ø@˜×+Ñ+×4Ñ4°T·Z±ZÑ@Ø+/¯9©9×+?Ò+?™ÀTôˆDŒJð ‰Dä‘7Ñ&Ó(ˆDÜ×+Ñ+¨D¯J©J¸¿	¹	À$¹ÔHØˆr!   c                ó:   — t        || j                  |dk(  |¬«      S )aŽ  Create a ClassificationDataset instance given an image path and mode.

        Args:
            img_path (str): Path to the dataset images.
            mode (str, optional): Dataset mode ('train', 'val', or 'test').
            batch (Any, optional): Batch information (unused in this implementation).

        Returns:
            (ClassificationDataset): Dataset for the specified mode.
        Útrain)Úrootr1   ÚaugmentÚprefix)r   r1   )r   Úimg_pathÚmodeÚbatchs       r    Úbuild_datasetz#ClassificationTrainer.build_datasetq   s   € ô %¨(¸¿¹ÈDÐT[ÉOÐdhÔiÐir!   c                ól  — t        |«      5  | j                  ||«      }ddd«       | j                  j                  dd«      }t	        j
                  j                  «      }|r–||kD  r‘|j
                  j                  |d }t	        |j                  «      }	|j                  D �
cg c]  }
|
d   |k  sŒ|
‘Œ c}
|_        |	t	        |j                  «      z
  }t        j                  |› d|› d|› d|› d|› �	«       t        ||| j                  j                  || j                  j                  ¬	«      }|d
k7  rkt        | j                  «      r1|j                   j"                  | j                  j$                  _        |S |j                   j"                  | j                  _        |S # 1 sw Y   �Œ‡xY wc c}
w )aÇ  Return PyTorch DataLoader with transforms to preprocess images.

        Args:
            dataset_path (str): Path to the dataset.
            batch_size (int, optional): Number of images per batch.
            rank (int, optional): Process rank for distributed training.
            mode (str, optional): 'train', 'val', or 'test' mode.

        Returns:
            (torch.utils.data.DataLoader): DataLoader for the specified dataset and mode.
        Nr)   r   é   z split has z classes but model expects z. Skipping z samples from extra classes: )ÚrankÚ	drop_lastrI   )r   rP   r$   r   ÚlenÚbaseÚclassesÚsamplesr   Úwarningr   r1   ÚworkersÚcompiler   r%   ÚdatasetÚtorch_transformsÚmoduleÚ
transforms)r   Údataset_pathÚ
batch_sizerS   rN   r\   r)   Ú
dataset_ncÚextra_classesÚoriginal_countÚsÚskippedÚloaders                r    Úget_dataloaderz$ClassificationTrainer.get_dataloader~   sv  € ô *¨$Ó/ñ 	=Ø×(Ñ(¨°tÓ<ˆG÷	=ð �Y‰Y�]‰]˜4 Ó#ˆÜ˜Ÿ™×-Ñ-Ó.ˆ
Ù�*˜r’/Ø#ŸL™L×0Ñ0°°Ð5ˆMÜ  §¡Ó1ˆNØ*1¯/©/ÖG Q¸Q¸q¹TÀB»YšqÒGˆGŒOØ$¤s¨7¯?©?Ó';Ñ;ˆGÜ�N‰NØ�&˜ J <Ð/JÈ2È$ð OØ#˜9Ð$AÀ-ÀðRôô
 " '¨:°t·y±y×7HÑ7HÈtÐ_c×_hÑ_h×_pÑ_pÔqˆà�7Š?Ü˜4Ÿ:™:Ô&Ø/5¯~©~×/NÑ/N�—
‘
×!Ñ!Ô,ð ˆð )/¯©×(GÑ(G�—
‘
Ô%Øˆ÷/	=ñ 	=üò Hs   ŒF$Â%F1Â3F1Æ$F.c                óî   — |d   j                  | j                  | j                  j                  dk(  ¬«      |d<   |d   j                  | j                  | j                  j                  dk(  ¬«      |d<   |S )z)Preprocess a batch of images and classes.ÚimgÚcuda)Únon_blockingÚcls)ÚtoÚdeviceÚtype)r   rO   s     r    Úpreprocess_batchz&ClassificationTrainer.preprocess_batch£   sc   € à˜U‘|—‘ t§{¡{ÀÇÁ×AQÑAQÐU[ÑA[�Ó\ˆˆe‰Ø˜U‘|—‘ t§{¡{ÀÇÁ×AQÑAQÐU[ÑA[�Ó\ˆˆe‰Øˆr!   c                ój   — dddt        | j                  «      z   z  z   ddg| j                  ¢d‘d‘­z  S )z4Return a formatted string showing training progress.ú
z%11sé   ÚEpochÚGPU_memÚ	InstancesÚSize)rU   Ú
loss_namesr&   s    r    Úprogress_stringz%ClassificationTrainer.progress_string©   sT   € à�v ¤S¨¯©Ó%9Ñ!9Ñ:Ñ:ØØð?
ð �_‰_ð?
ð ð	?
ð
 ñ?
ñ 
ð 	
r!   c                óº   — dg| _         t        j                  j                  | j                  | j
                  t        | j                  «      | j                  ¬«      S )z=Return an instance of ClassificationValidator for validation.Úloss)r1   r   )	ry   r	   r   ÚClassificationValidatorÚtest_loaderÚsave_dirr   r1   Ú	callbacksr&   s    r    Úget_validatorz#ClassificationTrainer.get_validator³   sF   € à!˜(ˆŒÜ�}‰}×4Ñ4Ø×Ñ˜dŸm™m´$°t·y±y³/ÈdÏnÉnð 5ó 
ð 	
r!   c                ó¦   — | j                   D �cg c]	  }|› d|› �‘Œ }}|€|S t        t        |«      d«      g}t        t	        ||«      «      S c c}w )a]  Return a loss dict with labeled training loss items tensor.

        Args:
            loss_items (torch.Tensor, optional): Loss tensor items.
            prefix (str, optional): Prefix to prepend to loss names.

        Returns:
            (dict | list): Dictionary of labeled loss items if loss_items is provided, otherwise list of keys.
        ú/é   )ry   ÚroundÚfloatÚdictÚzip)r   Ú
loss_itemsrL   ÚxÚkeyss        r    Úlabel_loss_itemsz&ClassificationTrainer.label_loss_itemsº   s[   € ð *.¯©Ö9 A�6�(˜!˜A˜3’Ð9ˆÐ9ØÐØˆKÜœE *Ó-¨qÓ1Ð2ˆ
Ü”C˜˜jÓ)Ó*Ð*ùò	 :s   �Ac                ó¦   — t        j                  |d   j                  d   «      |d<   t        || j                  d|› d�z  | j
                  ¬«       y)zßPlot training samples with their annotations.

        Args:
            batch (dict[str, torch.Tensor]): Batch containing images and class labels.
            ni (int): Batch index used for naming the output file.
        rj   r   Ú	batch_idxÚtrain_batchz.jpg)ÚlabelsÚfnameÚon_plotN)r5   ÚarangeÚshaper   r   r’   )r   rO   Únis      r    Úplot_training_samplesz+ClassificationTrainer.plot_training_samplesÊ   sM   € ô #Ÿ\™\¨%°©,×*<Ñ*<¸QÑ*?Ó@ˆˆkÑÜØØ—-‘- K°¨t°4Ð"8Ñ8Ø—L‘Lö	
r!   )r   zdict[str, Any] | Noner   zdict | None)NNT)r-   Úbool)rI   N)rM   rB   rN   rB   )é   r   rI   )r`   rB   ra   ÚintrS   r™   rN   rB   )rO   údict[str, torch.Tensor]Úreturnrš   )r›   rB   )NrI   )r‰   ztorch.Tensor | NonerL   rB   )rO   rš   r•   r™   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r'   r>   rE   rP   rh   rq   rz   r�   rŒ   r–   Ú__classcell__)r   s   @r    r   r      sL   ø„ ñð@ 'È4Ðkoö 5ò.ôô0ô$jô#óJó
ò
ô+÷ 
r!   r   )Ú
__future__r   r   Útypingr   r5   Úultralytics.datar   r   Úultralytics.engine.trainerr   Úultralytics.modelsr	   Úultralytics.nn.tasksr
   Úultralytics.utilsr   r   r   Úultralytics.utils.plottingr   Úultralytics.utils.torch_utilsr   r   r   © r!   r    ú<module>r«      s9   ðõ #å Ý ã ç DÝ 2Ý #Ý 4ß 7Ñ 7Ý 2ß SôC
˜Kõ C
r!   