Ë
    Fêñi+  ã                  ó¦   — d dl mZ d dlmZ d dlmZ d dlmZ d dlmZ d dl	m
Z
 d dlmZmZ d dlmZ  G d	„ d
ej                   j"                  «      Zy)é    )Úannotations)Úcopy)ÚPath)ÚAny)Úyolo)Ú	PoseModel)ÚDEFAULT_CFGÚRANK)Úunwrap_modelc                  óf   ‡ — e Zd ZdZeddfdˆ fd„Z	 	 	 d		 	 	 	 	 	 	 d
d„Zˆ fd„Zd„ Zdˆ fd„Z	ˆ xZ
S )ÚPoseTrainera¶  A class extending the DetectionTrainer class for training YOLO pose estimation models.

    This trainer specializes in handling pose estimation tasks, managing model training, validation, and visualization
    of pose keypoints alongside bounding boxes.

    Attributes:
        args (dict): Configuration arguments for training.
        model (PoseModel): The pose estimation model being trained.
        data (dict): Dataset configuration including keypoint shape information.
        loss_names (tuple): Names of the loss components used in training.

    Methods:
        get_model: Retrieve a pose estimation model with specified configuration.
        set_model_attributes: Set keypoints shape attribute on the model.
        get_validator: Create a validator instance for model evaluation.
        plot_training_samples: Visualize training samples with keypoints.
        get_dataset: Retrieve the dataset and ensure it contains required kpt_shape key.

    Examples:
        >>> from ultralytics.models.yolo.pose import PoseTrainer
        >>> args = dict(model="yolo26n-pose.pt", data="coco8-pose.yaml", epochs=3)
        >>> trainer = PoseTrainer(overrides=args)
        >>> trainer.train()
    Nc                ó:   •— |€i }d|d<   t         ‰| �  |||«       y)aw  Initialize a PoseTrainer object for training YOLO pose estimation 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.

        Notes:
            This trainer will automatically set the task to 'pose' regardless of what is provided in overrides.
            A warning is issued when using Apple MPS device due to known bugs with pose models.
        NÚposeÚtask)ÚsuperÚ__init__)ÚselfÚcfgÚ	overridesÚ
_callbacksÚ	__class__s       €úd/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/yolo/pose/train.pyr   zPoseTrainer.__init__)   s+   ø€ ð ÐØˆIØ"ˆ	�&ÑÜ‰Ñ˜˜i¨Õ4ó    c                ó°   — t        || j                  d   | j                  d   | j                  d   |xr	 t        dk(  ¬«      }|r|j                  |«       |S )a“  Get pose estimation model with specified configuration and weights.

        Args:
            cfg (str | Path | dict, optional): Model configuration file path or dictionary.
            weights (str | Path, optional): Path to the model weights file.
            verbose (bool): Whether to display model information.

        Returns:
            (PoseModel): Initialized pose estimation model.
        ÚncÚchannelsÚ	kpt_shapeéÿÿÿÿ)r   ÚchÚdata_kpt_shapeÚverbose)r   Údatar
   Úload)r   r   Úweightsr!   Úmodels        r   Ú	get_modelzPoseTrainer.get_model:   sV   € ô  ØØ�y‰y˜‰Ø�y‰y˜Ñ$ØŸ9™9 [Ñ1ØÒ*¤¨¡
ô
ˆñ Ø�J‰J�wÔàˆr   c           	     ó�  •— t         ‰| �  «        | j                  d   | j                  _        | j                  j                  d«      }|sft        t        t        t        | j                  j                  d   «      «      «      }t        | j                  j                  «      D �ci c]  }||“Œ }}|| j                  _        yc c}w )z+Set keypoints shape attribute of PoseModel.r   Ú	kpt_namesr   N)r   Úset_model_attributesr"   r%   r   ÚgetÚlistÚmapÚstrÚranger   r(   )r   r(   ÚnamesÚir   s       €r   r)   z PoseTrainer.set_model_attributesV   s“   ø€ ä‰Ñ$Ô&Ø#Ÿy™y¨Ñ5ˆ�
‰
ÔØ—I‘I—M‘M +Ó.ˆ	ÙÜœœS¤%¨¯
©
×(<Ñ(<¸QÑ(?Ó"@ÓAÓBˆEÜ+0°·±·±Ó+?Ö@ a˜˜E™Ð@ˆIÐ@Ø(ˆ�
‰
Õùò As   Â%
Cc                ó<  — d| _         t        t        | j                  «      j                  d   dd«      �| xj                   dz  c_         t        j
                  j                  | j                  | j                  t        | j                  «      | j                  ¬«      S )z=Return an instance of the PoseValidator class for validation.)Úbox_lossÚ	pose_lossÚ	kobj_lossÚcls_lossÚdfl_lossr   Ú
flow_modelN)Úrle_loss)Úsave_dirÚargsr   )Ú
loss_namesÚgetattrr   r%   r   r   ÚPoseValidatorÚtest_loaderr9   r   r:   Ú	callbacks)r   s    r   Úget_validatorzPoseTrainer.get_validator`   sx   € àVˆŒÜ”< §
¡
Ó+×1Ñ1°"Ñ5°|ÀTÓJÐVØ�OŠO˜}Ñ,�OÜ�y‰y×&Ñ&Ø×Ñ t§}¡}¼4ÀÇ	Á	»?ÐW[×WeÑWeð 'ó 
ð 	
r   c                ór   •— t         ‰| �  «       }d|vr#t        d| j                  j                  › d�«      ‚|S )a&  Retrieve the dataset and ensure it contains the required `kpt_shape` key.

        Returns:
            (dict): A dictionary containing the training/validation/test dataset and category names.

        Raises:
            KeyError: If the `kpt_shape` key is not present in the dataset.
        r   zNo `kpt_shape` in the z1. See https://docs.ultralytics.com/datasets/pose/)r   Úget_datasetÚKeyErrorr:   r"   )r   r"   r   s     €r   rB   zPoseTrainer.get_dataseti   s>   ø€ ô ‰wÑ"Ó$ˆØ˜dÑ"ÜÐ3°D·I±I·N±NÐ3CÐCtÐuÓvÐvØˆr   )r   zdict[str, Any] | Noner   zdict | None)NNT)r   z"str | Path | dict[str, Any] | Noner$   zstr | Path | Noner!   ÚboolÚreturnr   )rE   zdict[str, Any])Ú__name__Ú
__module__Ú__qualname__Ú__doc__r	   r   r&   r)   r@   rB   Ú__classcell__)r   s   @r   r   r      sa   ø„ ñð2 'È4Ðkoö 5ð& 37Ø%)Øð	à/ðð #ðð ð	ð
 
óô8)ò
÷ñ r   r   N)Ú
__future__r   r   Úpathlibr   Útypingr   Úultralytics.modelsr   Úultralytics.nn.tasksr   Úultralytics.utilsr	   r
   Úultralytics.utils.torch_utilsr   ÚdetectÚDetectionTrainerr   © r   r   ú<module>rU      s7   ðõ #å Ý Ý å #Ý *ß /Ý 6ôf�$—+‘+×.Ñ.õ fr   