Ë
    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 d dl	m
Z
mZmZ ddlmZ  G d„ d	e«      Z G d
„ de«      Zy)é    )Úannotations)ÚPath)ÚAnyN)Ú	IS_JETSONÚLOGGERÚ	is_jetsoné   )ÚBaseBackendc                  óh   ‡ — e Zd ZdZ	 	 	 d	 	 	 	 	 	 	 	 	 dˆ fd„Zdd„Z	 d	 	 	 	 	 	 	 	 	 	 	 d	d„Zˆ xZS )
ÚPyTorchBackendzÿPyTorch inference backend for native model execution.

    Loads and runs inference with native PyTorch models (.pt checkpoint files) or pre-loaded nn.Module
    instances. Supports model layer fusion, FP16 precision, and NVIDIA Jetson compatibility.
    c                óD   •— || _         || _        t        ‰| �  |||«       y)aã  Initialize the PyTorch backend.

        Args:
            weight (str | Path | nn.Module): Path to the .pt model file or a pre-loaded nn.Module instance.
            device (torch.device): Device to run inference on (e.g., 'cpu', 'cuda:0').
            fp16 (bool): Whether to use FP16 half-precision inference.
            fuse (bool): Whether to fuse Conv2D + BatchNorm layers for optimization.
            verbose (bool): Whether to print verbose model loading messages.
        N)ÚfuseÚverboseÚsuperÚ__init__)ÚselfÚweightÚdeviceÚfp16r   r   Ú	__class__s         €úa/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/nn/backends/pytorch.pyr   zPyTorchBackend.__init__   s$   ø€ ð" ˆŒ	ØˆŒÜ‰Ñ˜ ¨Õ.ó    c                óØ  — ddl m} t        |t        j                  j
                  «      r}| j                  rUt        |d«      rIt        r't        d¬«      r|j                  | j                  «      }|j                  | j                  ¬«      }|j                  | j                  «      }n" ||| j                  | j                  ¬«      \  }}t        |d«      r|j                  | _        t        |d	«      r-t        t        |j                   j                  «       «      d
«      nd
| _        t        |d«      r|j"                  j$                  nt'        |di «      | _        t        |d«      r|j(                  j+                  dd«      nd| _        | j.                  r|j1                  «       n|j3                  «        |j5                  «       D ]	  }d|_        Œ || _        t'        |dd«      | _        y)z¹Load a PyTorch model from a checkpoint file or nn.Module instance.

        Args:
            weight (str | torch.nn.Module): Path to the .pt checkpoint or a pre-loaded module.
        r   )Úload_checkpointr   é   )Újetpack)r   )r   r   Ú	kpt_shapeÚstrideé    ÚmoduleÚnamesÚyamlÚchannelsé   FÚend2endN)Úultralytics.nn.tasksr   Ú
isinstanceÚtorchÚnnÚModuler   Úhasattrr   r   Útor   r   r   ÚmaxÚintr   r    r!   Úgetattrr"   Úgetr#   r   ÚhalfÚfloatÚ
parametersÚrequires_gradÚmodelr%   )r   r   r   r5   Ú_Úps         r   Ú
load_modelzPyTorchBackend.load_model,   sb  € õ 	9ä�fœeŸh™hŸo™oÔ.Ø�yŠyœW V¨VÔ4Ý¤°1Õ!5Ø#ŸY™Y t§{¡{Ó3�FØŸ™¨T¯\©\˜Ó:�Ø—I‘I˜dŸk™kÓ*‰Eá& v°d·k±kÈÏ	É	ÔR‰HˆE�1ô �5˜+Ô&Ø"Ÿ_™_ˆDŒNÜ:AÀ%ÈÔ:R”cœ#˜eŸl™l×.Ñ.Ó0Ó1°2Ô6ÐXZˆŒÜ+2°5¸(Ô+C�U—\‘\×'Ò'ÌÐQVÐX_ÐacÓIdˆŒ
Ü9@ÀÈÔ9O˜Ÿ
™
Ÿ™ z°1Ô5ÐUVˆŒØŸ	š	ˆ�
‰
Œ u§{¡{£}øà×!Ñ!Ó#ò 	$ˆAØ#ˆA�Oð	$ð ˆŒ
Ü˜u i°Ó7ˆ�r   c                ó0   —  | j                   |f|||dœ|¤ŽS )ay  Run native PyTorch inference with support for augmentation, visualization, and embeddings.

        Args:
            im (torch.Tensor): Input image tensor in BCHW format, normalized to [0, 1].
            augment (bool): Whether to apply test-time augmentation.
            visualize (bool): Whether to visualize intermediate feature maps.
            embed (list | None): List of layer indices to extract embeddings from, or None.
            **kwargs (Any): Additional keyword arguments passed to the model forward method.

        Returns:
            (torch.Tensor | list[torch.Tensor]): Model predictions as tensor(s).
        )ÚaugmentÚ	visualizeÚembed©r5   )r   Úimr:   r;   r<   Úkwargss         r   ÚforwardzPyTorchBackend.forwardK   s$   € ð ˆt�z‰z˜"ÐZ g¸È%ÑZÐSYÑZÐZr   )FTT)
r   zstr | Path | nn.Moduler   útorch.devicer   Úboolr   rB   r   rB   )r   zstr | torch.nn.ModuleÚreturnÚNone)FFN)r>   útorch.Tensorr:   rB   r;   rB   r<   zlist | Noner?   r   rC   ú!torch.Tensor | list[torch.Tensor]©Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r8   r@   Ú__classcell__©r   s   @r   r   r      s‘   ø„ ñð ØØð/à&ð/ð ð/ð ð	/ð
 ð/ð õ/ó*8ð@ fjð[Øð[Ø)-ð[ØBFð[ØWbð[Øuxð[à	*÷[r   r   c                  ó6   ‡ — e Zd ZdZddˆ fd„Zdd„Zdd„Zˆ xZS )	ÚTorchScriptBackenda  PyTorch TorchScript inference backend for serialized model execution.

    Loads and runs inference with TorchScript models (.torchscript files) created via torch.jit.trace or
    torch.jit.script. Supports FP16 precision and embedded metadata extraction.
    c                ó(   •— t         ‰| �  |||«       y)a  Initialize the TorchScript backend.

        Args:
            weight (str | Path): Path to the .torchscript model file.
            device (torch.device): Device to run inference on (e.g., 'cpu', 'cuda:0').
            fp16 (bool): Whether to use FP16 half-precision inference.
        N)r   r   )r   r   r   r   r   s       €r   r   zTorchScriptBackend.__init__d   s   ø€ ô 	‰Ñ˜ ¨Õ.r   c                óˆ  — ddl }ddl}t        j                  d|› d�«       ddi}t        j
                  j                  ||| j                  ¬«      | _        | j                  r| j                  j                  «       n| j                  j                  «        |d   r'| j                  |j                  |d   d„ ¬	«      «       yy)
z©Load a TorchScript model from a .torchscript file with optional embedded metadata.

        Args:
            weight (str): Path to the .torchscript model file.
        r   NzLoading z for TorchScript inference...z
config.txtÚ )Ú_extra_filesÚmap_locationc                ó4   — t        | j                  «       «      S )N)ÚdictÚitems)Úxs    r   ú<lambda>z/TorchScriptBackend.load_model.<locals>.<lambda>~   s   € Ô\`Ðab×ahÑahÓajÓ\k€ r   )Úobject_hook)ÚjsonÚtorchvisionr   Úinfor(   ÚjitÚloadr   r5   r   r1   r2   Úapply_metadataÚloads)r   r   r[   r\   Úextra_filess        r   r8   zTorchScriptBackend.load_modeln   s�   € ó 	ãä�‰�h˜v˜hÐ&CÐDÔEØ# RÐ(ˆÜ—Y‘Y—^‘^ F¸ÐSW×S^ÑS^�^Ó_ˆŒ
Ø!ŸYšYˆ�
‰
�‰Ô¨D¯J©J×,<Ñ,<Ó,>øà�|Ò$Ø×Ñ §
¡
¨;°|Ñ+DÑRk 
Ó lÕmð %r   c                ó$   — | j                  |«      S )zíRun TorchScript inference.

        Args:
            im (torch.Tensor): Input image tensor in BCHW format, normalized to [0, 1].

        Returns:
            (torch.Tensor | list[torch.Tensor]): Model predictions as tensor(s).
        r=   )r   r>   s     r   r@   zTorchScriptBackend.forward€   s   € ð �z‰z˜"‹~Ðr   )F)r   z
str | Pathr   rA   r   rB   )r   ÚstrrC   rD   )r>   rE   rC   rF   rG   rM   s   @r   rO   rO   ]   s   ø„ ñö/ón÷$	r   rO   )Ú
__future__r   Úpathlibr   Útypingr   r(   Útorch.nnr)   Úultralytics.utilsr   r   r   Úbaser
   r   rO   © r   r   ú<module>rl      s<   ðõ #å Ý ã Ý ç :Ñ :å ôJ[�[ô J[ôZ,˜õ ,r   