Ë
    Fêñi-.  ã                  óÞ   — d dl mZ d dlZd dlZd dlmZ d dlmZ d dl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 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 d dlmZm Z   G d„ de«      Z!y)é    )ÚannotationsN)Úcopy)ÚAny)Úbuild_dataloaderÚbuild_yolo_dataset)ÚBaseTrainer)Úyolo)ÚDetectionModel)ÚDEFAULT_CFGÚLOGGERÚRANK)Úoverride_configs)Úplot_imagesÚplot_labels)Útorch_distributed_zero_firstÚunwrap_modelc                  óŒ   ‡ — e Zd ZdZeddfdˆ fd„Zddd„Zddd„Zdd„Zd„ Z	d„ Z
ddd	„Zd
„ Zddd„Zd„ Zdd„Zd„ Zˆ fd„Zˆ xZS )ÚDetectionTrainera©  A class extending the BaseTrainer class for training based on a detection model.

    This trainer specializes in object detection tasks, handling the specific requirements for training YOLO models for
    object detection including dataset building, data loading, preprocessing, and model configuration.

    Attributes:
        model (DetectionModel): The YOLO detection model being trained.
        data (dict): Dictionary containing dataset information including class names and number of classes.
        loss_names (tuple): Names of the loss components used in training (box_loss, cls_loss, dfl_loss).

    Methods:
        build_dataset: Build YOLO dataset for training or validation.
        get_dataloader: Construct and return dataloader for the specified mode.
        preprocess_batch: Preprocess a batch of images by scaling and converting to float.
        set_model_attributes: Set model attributes based on dataset information.
        get_model: Return a YOLO detection model.
        get_validator: Return a validator for model evaluation.
        label_loss_items: Return a loss dictionary with labeled training loss items.
        progress_string: Return a formatted string of training progress.
        plot_training_samples: Plot training samples with their annotations.
        plot_training_labels: Create a labeled training plot of the YOLO model.
        auto_batch: Calculate optimal batch size based on model memory requirements.

    Examples:
        >>> from ultralytics.models.yolo.detect import DetectionTrainer
        >>> args = dict(model="yolo26n.pt", data="coco8.yaml", epochs=3)
        >>> trainer = DetectionTrainer(overrides=args)
        >>> trainer.train()
    Nc                ó(   •— t         ‰| �  |||«       y)a�  Initialize a DetectionTrainer object for training YOLO object detection models.

        Args:
            cfg (dict, optional): Default configuration dictionary containing training parameters.
            overrides (dict, optional): Dictionary of parameter overrides for the default configuration.
            _callbacks (dict, optional): Dictionary of callback functions to be executed during training.
        N)ÚsuperÚ__init__)ÚselfÚcfgÚ	overridesÚ
_callbacksÚ	__class__s       €úf/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/yolo/detect/train.pyr   zDetectionTrainer.__init__7   s   ø€ ô 	‰Ñ˜˜i¨Õ4ó    c           	     óÔ   — t        t        t        | j                  «      j                  j                  «       «      d«      }t        | j                  ||| j                  ||dk(  |¬«      S )a¬  Build YOLO Dataset for training or validation.

        Args:
            img_path (str): Path to the folder containing images.
            mode (str): 'train' mode or 'val' mode, users are able to customize different augmentations for each mode.
            batch (int, optional): Size of batches, this is for 'rect' mode.

        Returns:
            (Dataset): YOLO dataset object configured for the specified mode.
        é    Úval)ÚmodeÚrectÚstride)ÚmaxÚintr   Úmodelr$   r   ÚargsÚdata)r   Úimg_pathr"   ÚbatchÚgss        r   Úbuild_datasetzDetectionTrainer.build_datasetA   sT   € ô ””\ $§*¡*Ó-×4Ñ4×8Ñ8Ó:Ó;¸RÓ@ˆÜ! $§)¡)¨X°u¸d¿i¹iÈdÐY]ÐafÑYfÐoqÔrÐrr   c           	     óö  — |dv sJ d|› d�«       ‚t        |«      5  | j                  |||«      }ddd«       |dk(  }t        dd«      rH|rFt        j                  |j
                  |j
                  d   k(  «      st        j                  d	«       d}t        |||dk(  r| j                  j                  n| j                  j                  d
z  ||| j                  j                  xr |dk(  ¬«      S # 1 sw Y   ŒÁxY w)až  Construct and return dataloader for the specified mode.

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

        Returns:
            (DataLoader): PyTorch dataloader object.
        >   r!   Útrainz#Mode must be 'train' or 'val', not ú.Nr/   r#   Fr   zJ'rect=True' is incompatible with DataLoader shuffle, setting shuffle=Falseé   )r+   ÚworkersÚshuffleÚrankÚ	drop_last)r   r-   ÚgetattrÚnpÚallÚbatch_shapesr   Úwarningr   r(   r2   Úcompile)r   Údataset_pathÚ
batch_sizer4   r"   Údatasetr3   s          r   Úget_dataloaderzDetectionTrainer.get_dataloaderO   sø   € ð Ð'Ñ'ÐVÐ+NÈtÈfÐTUÐ)VÓVÐ'Ü)¨$Ó/ñ 	IØ×(Ñ(¨°t¸ZÓHˆG÷	Ià˜'‘/ˆÜ�7˜F EÔ*©w¼r¿v¹vÀg×FZÑFZÐ^e×^rÑ^rÐstÑ^uÑFuÔ?vÜ�N‰NÐgÔhØˆGÜØØØ)-°ª�D—I‘I×%Ò%¸d¿i¹i×>OÑ>OÐRSÑ>SØØØ—i‘i×'Ñ'Ò;¨D°G©Oô
ð 	
÷	Ið 	Iús   ›C/Ã/C8c                óÒ  — |j                  «       D ]W  \  }}t        |t        j                  «      sŒ!|j	                  | j
                  | j
                  j                  dk(  ¬«      ||<   ŒY |d   j                  «       dz  |d<   | j                  j                  dkD  �rD|d   }t        j                  t        | j                  j                  d| j                  j                  z
  z  «      t        | j                  j                  d| j                  j                  z   z  | j                  z   «      «      | j                  z  | j                  z  }|t        |j                   dd «      z  }|d	k7  ro|j                   dd D �cg c]4  }t#        j$                  ||z  | j                  z  «      | j                  z  ‘Œ6 }}t&        j(                  j+                  ||d
d¬«      }||d<   |S c c}w )z÷Preprocess a batch of images by scaling and converting to float.

        Args:
            batch (dict): Dictionary containing batch data with 'img' tensor.

        Returns:
            (dict): Preprocessed batch with normalized images.
        Úcuda)Únon_blockingÚimgéÿ   ç        ç      ð?r1   Né   ÚbilinearF)Úsizer"   Úalign_corners)ÚitemsÚ
isinstanceÚtorchÚTensorÚtoÚdeviceÚtypeÚfloatr(   Úmulti_scaleÚrandomÚ	randranger&   Úimgszr$   r%   ÚshapeÚmathÚceilÚnnÚ
functionalÚinterpolate)	r   r+   ÚkÚvÚimgsÚszÚsfÚxÚnss	            r   Úpreprocess_batchz!DetectionTrainer.preprocess_batchk   s¨  € ð —K‘K“Mò 	V‰DˆAˆqÜ˜!œUŸ\™\Õ*ØŸ4™4 §¡¸$¿+¹+×:JÑ:JÈfÑ:T˜4ÓU��a’ð	Vð ˜U‘|×)Ñ)Ó+¨cÑ1ˆˆe‰Ø�9‰9× Ñ  3Ó&Ø˜‘<ˆDä× Ñ Ü˜Ÿ	™	Ÿ™¨3°·±×1FÑ1FÑ+FÑGÓHÜ˜Ÿ	™	Ÿ™¨3°·±×1FÑ1FÑ+FÑGÈ$Ï+É+ÑUÓVóð —;‘;ñ	ð
 —+‘+ñð ð ”c˜$Ÿ*™* Q R˜.Ó)Ñ)ˆBØ�QŠwàKOÏ:É:ÐVWÐVXÈ>öØFG”D—I‘I˜a "™f t§{¡{Ñ2Ó3°d·k±kÓAð�ð ô —}‘}×0Ñ0°¸BÀZÐ_dÐ0Óe�ØˆE�%‰LØˆùòs   Å?9G$c                ó@  — | j                   d   | j                  _        | j                   d   | j                  _        | j                  | j                  _        t        | j                  d«      r1| j                  j                  | j                  j                  ¬«       yy)z2Set model attributes based on dataset information.ÚncÚnamesÚend2end)Úmax_detN)r)   r'   rf   rg   r(   r6   Úset_head_attrri   ©r   s    r   Úset_model_attributesz%DetectionTrainer.set_model_attributes‹   sm   € ð Ÿ	™	 $™ˆ�
‰
ŒØŸ9™9 WÑ-ˆ�
‰
ÔØŸ)™)ˆ�
‰
ŒÜ�4—:‘:˜yÔ)Ø�J‰J×$Ñ$¨T¯Y©Y×->Ñ->Ð$Õ?ð *r   c                ó¦  — d| j                   j                  cxk  rdk  sJ d«       ‚ J d«       ‚| j                   j                  dk(  ryt        j                  | j                  j
                  j                  D �cg c]  }|d   j                  «       ‘Œ c}d«      }t        j                  |j                  t        «      | j                  d   ¬«      j                  t        j                  «      }t        j                  |dk(  d|«      }d|z  | j                   j                  z  }||j                  «       z  }t        j                   |«      j#                  | j$                  «      | j&                  _        t+        j,                  d	| j&                  j(                  j/                  «       j1                  «       j3                  d
«      › �«       yc c}w )a8  Compute and set class weights for handling class imbalance.

        Class weights are computed based on inverse class frequency in the training dataset,
        raised to the power of cls_pw (0 < cls_pw <= 1 dampens, cls_pw > 1 amplifies).
        Final weights are normalized so their mean equals 1.0.
        r   rF   z"cls_pw must be in the range [0, 1]rE   NÚclsrf   )Ú	minlengthzClass weights: é   )r(   Úcls_pwr7   ÚconcatenateÚtrain_loaderr>   ÚlabelsÚflattenÚbincountÚastyper&   r)   Úfloat32ÚwhereÚmeanrM   Ú
from_numpyrO   rP   r'   Úclass_weightsr   ÚinfoÚcpuÚnumpyÚround)r   ÚlbÚclassesÚclass_countsÚweightss        r   Úset_class_weightsz"DetectionTrainer.set_class_weights—   s^  € ð �D—I‘I×$Ñ$Ô+¨Ò+ÐQÐ-QÓQÑ+ÐQÐ-QÓQÐ+Ø�9‰9×Ñ˜sÒ"ØÜ—.‘.À×@QÑ@Q×@YÑ@Y×@`Ñ@`Ö!a¸" " U¡)×"3Ñ"3Õ"5Ò!aÐcdÓeˆÜ—{‘{ 7§>¡>´#Ó#6À$Ç)Á)ÈDÁ/ÔR×YÑYÔZ\×ZdÑZdÓeˆÜ—x‘x °Ñ 1°3¸ÓEˆà˜Ñ%¨$¯)©)×*:Ñ*:Ñ:ˆØ˜GŸL™L›NÑ*ˆÜ#(×#3Ñ#3°GÓ#<×#?Ñ#?ÀÇÁÓ#Lˆ�
‰
Ô Ü�‰�o d§j¡j×&>Ñ&>×&BÑ&BÓ&D×&JÑ&JÓ&L×&RÑ&RÐSTÓ&UÐ%VÐWÕXùò "bs   Á:Gc                ó”   — t        || j                  d   | j                  d   |xr	 t        dk(  ¬«      }|r|j                  |«       |S )a=  Return a YOLO detection model.

        Args:
            cfg (str, optional): Path to model configuration file.
            weights (str, optional): Path to model weights.
            verbose (bool): Whether to display model information.

        Returns:
            (DetectionModel): YOLO detection model.
        rf   Úchannelséÿÿÿÿ)rf   ÚchÚverbose)r
   r)   r   Úload)r   r   r„   rŠ   r'   s        r   Ú	get_modelzDetectionTrainer.get_modelª   sF   € ô ˜s t§y¡y°¡¸4¿9¹9ÀZÑ;PÐZaÒZpÔfjÐnpÑfpÔqˆÙØ�J‰J�wÔØˆr   c                ó¸   — d| _         t        j                  j                  | j                  | j
                  t        | j                  «      | j                  ¬«      S )z6Return a DetectionValidator for YOLO model validation.)Úbox_lossÚcls_lossÚdfl_loss)Úsave_dirr(   r   )	Ú
loss_namesr	   ÚdetectÚDetectionValidatorÚtest_loaderr‘   r   r(   Ú	callbacksrk   s    r   Úget_validatorzDetectionTrainer.get_validatorº   sG   € à<ˆŒÜ�{‰{×-Ñ-Ø×Ñ t§}¡}¼4ÀÇ	Á	»?ÐW[×WeÑWeð .ó 
ð 	
r   c                óÈ   — | j                   D �cg c]	  }|› d|› �‘Œ }}|�7|D �cg c]  }t        t        |«      d«      ‘Œ }}t        t	        ||«      «      S |S c c}w c c}w )a_  Return a loss dict with labeled training loss items tensor.

        Args:
            loss_items (list[float], optional): List of loss values.
            prefix (str): Prefix for keys in the returned dictionary.

        Returns:
            (dict | list): Dictionary of labeled loss items if loss_items is provided, otherwise list of keys.
        ú/é   )r’   r€   rR   ÚdictÚzip)r   Ú
loss_itemsÚprefixrb   Úkeyss        r   Úlabel_loss_itemsz!DetectionTrainer.label_loss_itemsÁ   sh   € ð *.¯©Ö9 A�6�(˜!˜A˜3’Ð9ˆÐ9ØÐ!Ø6@ÖA°œ%¤ a£¨!Õ,ÐAˆJÐAÜœ˜D *Ó-Ó.Ð.àˆKùò :ùâAs
   �A¥Ac                ój   — dddt        | j                  «      z   z  z   ddg| j                  ¢d‘d‘­z  S )z`Return a formatted string of training progress with epoch, GPU memory, loss, instances and size.ú
z%11sé   ÚEpochÚGPU_memÚ	InstancesÚSize)Úlenr’   rk   s    r   Úprogress_stringz DetectionTrainer.progress_stringÒ   sT   € à�v ¤S¨¯©Ó%9Ñ!9Ñ:Ñ:ØØð?
ð �_‰_ð?
ð ð	?
ð
 ñ?
ñ 
ð 	
r   c                ó^   — t        ||d   | j                  d|› d�z  | j                  ¬«       y)zÎPlot training samples with their annotations.

        Args:
            batch (dict[str, Any]): Dictionary containing batch data.
            ni (int): Batch index used for naming the output file.
        Úim_fileÚtrain_batchz.jpg)rt   ÚpathsÚfnameÚon_plotN)r   r‘   r¯   )r   r+   Únis      r   Úplot_training_samplesz&DetectionTrainer.plot_training_samplesÜ   s3   € ô 	ØØ˜	Ñ"Ø—-‘- K°¨t°4Ð"8Ñ8Ø—L‘Lö		
r   c                óª  — t        j                  | j                  j                  j                  D �cg c]  }|d   ‘Œ	 c}d«      }t        j                  | j                  j                  j                  D �cg c]  }|d   ‘Œ	 c}d«      }t        ||j                  «       | j                  d   | j                  | j                  ¬«       yc c}w c c}w )z1Create a labeled training plot of the YOLO model.Úbboxesr   rn   rg   )rg   r‘   r¯   N)
r7   rr   rs   r>   rt   r   Úsqueezer)   r‘   r¯   )r   r�   Úboxesrn   s       r   Úplot_training_labelsz%DetectionTrainer.plot_training_labelsê   s™   € ä—‘°t×7HÑ7H×7PÑ7P×7WÑ7WÖX°  8£ÒXÐZ[Ó\ˆÜ�n‰n°$×2CÑ2C×2KÑ2K×2RÑ2RÖS¨B˜b ›iÒSÐUVÓWˆÜ�E˜3Ÿ;™;›=°·	±	¸'Ñ0BÈTÏ]É]Ðdh×dpÑdpÖqùò  YùÚSs   ²CÁ7Cc                ó$  •— t        | j                  ddi¬«      5 | _        | j                  | j                  d   dd¬«      }ddd«       t	        d„ j
                  D «       «      d	z  }t        |«      }~t        ‰| �!  ||¬
«      S # 1 sw Y   ŒExY w)zƒGet optimal batch size by calculating memory occupation of model.

        Returns:
            (int): Optimal batch size.
        ÚcacheF)r   r/   é   )r"   r+   Nc              3  ó8   K  — | ]  }t        |d    «      –— Œ y­w)rn   N)r¨   )Ú.0Úlabels     r   ú	<genexpr>z.DetectionTrainer.auto_batch.<locals>.<genexpr>ø   s   è ø€ ÒN°œ#˜e E™l×+ÑNùs   ‚r£   )Údataset_size)	r   r(   r-   r)   r%   rt   r¨   r   Ú
auto_batch)r   Útrain_datasetÚmax_num_objÚnr   s       €r   r¿   zDetectionTrainer.auto_batchð   s”   ø€ ô ˜dŸi™i°G¸UÐ3CÔDð 	[ÈÌ	Ø ×.Ñ.¨t¯y©y¸Ñ/AÈÐWYÐ.ÓZˆM÷	[äÑN¸×9MÑ9MÔNÓNÐQRÑRˆÜ�ÓˆØÜ‰wÑ! +¸AÐ!Ó>Ð>÷	[ð 	[ús   ›'BÂB)r   zdict[str, Any] | Noner   zdict | None)r/   N)r*   Ústrr"   rÃ   r+   z
int | None)r¹   r   r/   )r<   rÃ   r=   r&   r4   r&   r"   rÃ   )r+   r›   Úreturnr›   )NNT)r   ú
str | Noner„   rÅ   rŠ   Úbool)Nr/   )r�   zlist[float] | Nonerž   rÃ   )r+   zdict[str, Any]r°   r&   rÄ   ÚNone)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r-   r?   rd   rl   r…   rŒ   r—   r    r©   r±   r¶   r¿   Ú__classcell__)r   s   @r   r   r      s]   ø„ ñð< 'È4Ðkoö 5ôsô
ó8ò@
@òYô&ò 
ôò"
ó
òr÷?ð ?r   r   )"Ú
__future__r   rX   rT   r   Útypingr   r   r7   rM   Útorch.nnrZ   Úultralytics.datar   r   Úultralytics.engine.trainerr   Úultralytics.modelsr	   Úultralytics.nn.tasksr
   Úultralytics.utilsr   r   r   Úultralytics.utils.patchesr   Úultralytics.utils.plottingr   r   Úultralytics.utils.torch_utilsr   r   r   © r   r   ú<module>rÙ      sH   ðõ #ã Û Ý Ý ã Û Ý ç AÝ 2Ý #Ý /ß 7Ñ 7Ý 6ß ?ß Tôc?�{õ c?r   