Ë
    Fêñij  ã                  óv   — 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m	Z	 d dl
mZ d dlmZmZ  G d„ d	e	«      Zy)
é    )ÚannotationsN)ÚImage)Úclassify_transforms)ÚBasePredictor)ÚResults)ÚDEFAULT_CFGÚopsc                  óB   ‡ — e Zd ZdZeddfdˆ fd„Zˆ fd„Zd„ Zd„ Zˆ xZ	S )ÚClassificationPredictora¬  A class extending the BasePredictor class for prediction based on a classification model.

    This predictor handles the specific requirements of classification models, including preprocessing images and
    postprocessing predictions to generate classification results.

    Attributes:
        args (dict): Configuration arguments for the predictor.

    Methods:
        preprocess: Convert input images to model-compatible format.
        postprocess: Process model predictions into Results objects.

    Examples:
        >>> from ultralytics.utils import ASSETS
        >>> from ultralytics.models.yolo.classify import ClassificationPredictor
        >>> args = dict(model="yolo26n-cls.pt", source=ASSETS)
        >>> predictor = ClassificationPredictor(overrides=args)
        >>> predictor.predict_cli()

    Notes:
        - Torchvision classification models can also be passed to the 'model' argument, i.e. model='resnet18'.
    Nc                óJ   •— t         ‰| �  |||«       d| j                  _        y)as  Initialize the ClassificationPredictor with the specified configuration and set task to 'classify'.

        This constructor initializes a ClassificationPredictor instance, which extends BasePredictor for classification
        tasks. It ensures the task is set to 'classify' regardless of input configuration.

        Args:
            cfg (dict): Default configuration dictionary containing prediction settings.
            overrides (dict, optional): Configuration overrides that take precedence over cfg.
            _callbacks (dict, optional): Dictionary of callback functions to be executed during prediction.
        ÚclassifyN)ÚsuperÚ__init__ÚargsÚtask)ÚselfÚcfgÚ	overridesÚ
_callbacksÚ	__class__s       €új/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/yolo/classify/predict.pyr   z ClassificationPredictor.__init__'   s!   ø€ ô 	‰Ñ˜˜i¨Ô4Ø#ˆ�	‰	�ó    c                ó&  •— t         ‰| �  |«       t        | j                  j                  d«      r„t        | j                  j                  j                  j                  d   d«      rM| j                  j                  j                  j                  d   j
                  t        | j                  «      k7  nd}|s| j                  j                  dk7  rt        | j                  «      | _        y| j                  j                  j                  | _        y)z9Set up source and inference mode and classify transforms.Ú
transformsr   ÚsizeFÚptN)
r   Úsetup_sourceÚhasattrÚmodelr   r   ÚmaxÚimgszÚformatr   )r   ÚsourceÚupdatedr   s      €r   r   z$ClassificationPredictor.setup_source5   sÌ   ø€ ä‰Ñ˜VÔ$ô �t—z‘z×'Ñ'¨Ô6¼7À4Ç:Á:×CSÑCS×C^ÑC^×CiÑCiÐjkÑClÐntÔ;uð �J‰J×Ñ×'Ñ'×2Ñ2°1Ñ5×:Ñ:¼cÀ$Ç*Á*»oÒMàð 	ñ 07¸$¿*¹*×:KÑ:KÈtÒ:SÔ §
¡
Ó+ð 	�ØY]×YcÑYc×YiÑYi×YtÑYtð 	�r   c                ó&  — t        |t        j                  «      sit        j                  |D �cg c]H  }| j	                  t        j                  t        j                  |t        j                  «      «      «      ‘ŒJ c}d¬«      }t        |t        j                  «      r|nt        j                  |«      j                  | j                  j                  «      }| j                  j                  r|j                  «       S |j!                  «       S c c}w )zVConvert input images to model-compatible tensor format with appropriate normalization.r   )Údim)Ú
isinstanceÚtorchÚTensorÚstackr   r   Ú	fromarrayÚcv2ÚcvtColorÚCOLOR_BGR2RGBÚ
from_numpyÚtor   ÚdeviceÚfp16ÚhalfÚfloat)r   ÚimgÚims      r   Ú
preprocessz"ClassificationPredictor.preprocessA   s³   € ä˜#œuŸ|™|Ô,Ü—+‘+ØadÖeÐ[]�—‘¤§¡´·±¸bÄ#×BSÑBSÓ1TÓ!UÕVÒeÐklôˆCô ! ¤e§l¡lÔ3‰s¼×9IÑ9IÈ#Ó9N×RÑRÐSW×S]ÑS]×SdÑSdÓeˆØ!ŸZ™ZŸ_š_ˆs�x‰x‹zÐ=°#·)±)³+Ð=ùò fs   ®ADc                óF  — t        |t        «      st        j                  |«      dddd…f   }t        |t        t        f«      r|d   n|}t        ||| j                  d   «      D ���cg c])  \  }}}t        ||| j                  j                  |¬«      ‘Œ+ c}}}S c c}}}w )aÄ  Process predictions to return Results objects with classification probabilities.

        Args:
            preds (torch.Tensor): Raw predictions from the model.
            img (torch.Tensor): Input images after preprocessing.
            orig_imgs (list[np.ndarray] | torch.Tensor): Original images before preprocessing.

        Returns:
            (list[Results]): List of Results objects containing classification results for each image.
        .Néÿÿÿÿr   )ÚpathÚnamesÚprobs)
r'   Úlistr	   Úconvert_torch2numpy_batchÚtupleÚzipÚbatchr   r   r;   )r   Úpredsr5   Ú	orig_imgsÚpredÚorig_imgÚimg_paths          r   Úpostprocessz#ClassificationPredictor.postprocessJ   s–   € ô ˜)¤TÔ*Ü×5Ñ5°iÓ@ÀÁdÈÀdÀÑKˆIä& u¬t´U¨mÔ<��a’À%ˆô -0°°yÀ$Ç*Á*ÈQÁ-Ó,P÷
ð 
á(��h ô �H 8°4·:±:×3CÑ3CÈ4ÖPô
ð 	
ùô 
s   Á).B)r   zdict | None)
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r7   rG   Ú__classcell__)r   s   @r   r   r      s&   ø„ ñð. '°$ÐRVö $ô

ò>ö
r   r   )Ú
__future__r   r,   r(   ÚPILr   Úultralytics.data.augmentr   Úultralytics.engine.predictorr   Úultralytics.engine.resultsr   Úultralytics.utilsr   r	   r   © r   r   ú<module>rT      s-   ðõ #ã 
Û Ý å 8Ý 6Ý .ß .ôM
˜mõ M
r   