Ë
    Fêñiœ  ã                  ój   — 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mZ  G d„ d	e«      Zy
)é    )Úannotations)Úcopy)ÚDetectionTrainer)ÚRTDETRDetectionModel)ÚRANKÚcolorstré   )ÚRTDETRDatasetÚRTDETRValidatorc                  ó*   — e Zd ZdZddd„Zdd	d„Zd„ Zy)
ÚRTDETRTrainera²  Trainer class for the RT-DETR model developed by Baidu for real-time object detection.

    This class extends the DetectionTrainer class for YOLO to adapt to the specific features and architecture of
    RT-DETR. The model leverages Vision Transformers and has capabilities like IoU-aware query selection and adaptable
    inference speed.

    Attributes:
        loss_names (tuple): Names of the loss components used for training.
        data (dict): Dataset configuration containing class count and other parameters.
        args (dict): Training arguments and hyperparameters.
        save_dir (Path): Directory to save training results.
        test_loader (DataLoader): DataLoader for validation/testing data.

    Methods:
        get_model: Initialize and return an RT-DETR model for object detection tasks.
        build_dataset: Build and return an RT-DETR dataset for training or validation.
        get_validator: Return a DetectionValidator suitable for RT-DETR model validation.

    Examples:
        >>> from ultralytics.models.rtdetr.train import RTDETRTrainer
        >>> args = dict(model="rtdetr-l.yaml", data="coco8.yaml", imgsz=640, epochs=3)
        >>> trainer = RTDETRTrainer(overrides=args)
        >>> trainer.train()

    Notes:
        - F.grid_sample used in RT-DETR does not support the `deterministic=True` argument.
        - AMP training can lead to NaN outputs and may produce errors during bipartite graph matching.
    Nc                ó”   — t        || j                  d   | j                  d   |xr	 t        dk(  ¬«      }|r|j                  |«       |S )aW  Initialize and return an RT-DETR model for object detection tasks.

        Args:
            cfg (dict, optional): Model configuration.
            weights (str, optional): Path to pre-trained model weights.
            verbose (bool): Verbose logging if True.

        Returns:
            (RTDETRDetectionModel): Initialized model.
        ÚncÚchannelséÿÿÿÿ)r   ÚchÚverbose)r   Údatar   Úload)ÚselfÚcfgÚweightsr   Úmodels        úa/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/rtdetr/train.pyÚ	get_modelzRTDETRTrainer.get_model,   sF   € ô % S¨T¯Y©Y°t©_ÀÇÁÈ:ÑAVÐ`gÒ`vÔlpÐtvÑlvÔwˆÙØ�J‰J�wÔØˆó    c                óf  — t        || j                  j                  ||dk(  | j                  d| j                  j                  xs d| j                  j                  xs dt        |› d�«      | j                  j                  | j                  |dk(  r| j                  j                  ¬«      S d¬«      S )as  Build and return an RT-DETR dataset for training or validation.

        Args:
            img_path (str): Path to the folder containing images.
            mode (str): Dataset mode, either 'train' or 'val'.
            batch (int, optional): Batch size for rectangle training.

        Returns:
            (RTDETRDataset): Dataset object for the specific mode.
        ÚtrainFNz: g      ð?)Úimg_pathÚimgszÚ
batch_sizeÚaugmentÚhypÚrectÚcacheÚ
single_clsÚprefixÚclassesr   Úfraction)	r
   Úargsr    r%   r&   r   r(   r   r)   )r   r   ÚmodeÚbatchs       r   Úbuild_datasetzRTDETRTrainer.build_dataset<   s›   € ô ØØ—)‘)—/‘/ØØ˜G‘OØ—	‘	ØØ—)‘)—/‘/Ò) TØ—y‘y×+Ñ+Ò4¨uÜ˜t˜f B˜KÓ(Ø—I‘I×%Ñ%Ø—‘Ø+/°7ª?�T—Y‘Y×'Ñ'ô
ð 	
ð ADô
ð 	
r   c                óz   — d| _         t        | j                  | j                  t	        | j
                  «      ¬«      S )z@Return an RTDETRValidator suitable for RT-DETR model validation.)Ú	giou_lossÚcls_lossÚl1_loss)Úsave_dirr*   )Ú
loss_namesr   Útest_loaderr2   r   r*   )r   s    r   Úget_validatorzRTDETRTrainer.get_validatorV   s-   € à<ˆŒÜ˜t×/Ñ/¸$¿-¹-ÌdÐSW×S\ÑS\ËoÔ^Ð^r   )NNT)r   zdict | Noner   z
str | Noner   Úbool)ÚvalN)r   Ústrr+   r8   r,   z
int | None)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r-   r5   © r   r   r   r      s   „ ñô:ô 
ó4_r   r   N)Ú
__future__r   r   Úultralytics.models.yolo.detectr   Úultralytics.nn.tasksr   Úultralytics.utilsr   r   r7   r
   r   r   r=   r   r   ú<module>rB      s*   ðõ #å å ;Ý 5ß ,ç /ôK_Ð$õ K_r   