# Ultralytics 🚀 AGPL-3.0 License - https://ultralytics.com/license

from __future__ import annotations

from pathlib import Path
from typing import Any

import torch
from torch import nn

from ultralytics.utils import IS_JETSON, LOGGER, is_jetson
from ultralytics.utils.torch_utils import TORCH_1_10, TORCH_2_1, unwrap_model

from .base import BaseBackend


class PyTorchBackend(BaseBackend):
    """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.
    """

    def __init__(
        self,
        weight: str | Path | nn.Module,
        device: torch.device,
        fp16: bool = False,
        fuse: bool = True,
        verbose: bool = True,
        end2end: bool | None = None,
    ):
        """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.
            end2end (bool, optional): Select the detection head before fusion; None preserves its current mode.
        """
        self.fuse = fuse
        self.verbose = verbose
        self.end2end_override = end2end
        super().__init__(weight, device, fp16)

    def load_model(self, weight: str | Path | nn.Module) -> None:
        """Load a PyTorch model from a checkpoint file or nn.Module instance.

        Args:
            weight (str | Path | nn.Module): Path to the .pt checkpoint or a pre-loaded module.
        """
        from ultralytics.nn.tasks import BaseModel, load_checkpoint

        model = weight if isinstance(weight, torch.nn.Module) else load_checkpoint(weight, device=self.device)[0]
        if self.end2end_override is not None and hasattr(model, "end2end"):
            model.end2end = self.end2end_override
        if self.fuse and hasattr(model, "fuse"):
            if IS_JETSON and is_jetson(jetpack=5):
                model = model.to(self.device)
            model = model.fuse(verbose=self.verbose) if isinstance(model, BaseModel) else model.fuse()
        model = model.to(self.device)

        # Extract model attributes
        if hasattr(model, "kpt_shape"):
            self.kpt_shape = model.kpt_shape
        self.stride = max(int(model.stride.max()), 32) if hasattr(model, "stride") else 32
        self.names = model.module.names if hasattr(model, "module") else getattr(model, "names", {})
        self.channels = model.yaml.get("channels", 3) if hasattr(model, "yaml") else 3
        model.half() if self.fp16 else model.float()

        for p in model.parameters():
            p.requires_grad = False

        self.model = model
        self.end2end = getattr(model, "end2end", False)
        self.base_model = isinstance(unwrap_model(model), BaseModel)

    def forward(
        self, im: torch.Tensor, augment: bool = False, embed: list | None = None, **kwargs: Any
    ) -> torch.Tensor | list[torch.Tensor]:
        """Run native PyTorch inference with support for augmentation and embeddings.

        Args:
            im (torch.Tensor): Input image tensor in BCHW format, normalized to [0, 1].
            augment (bool): Whether to apply test-time augmentation.
            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).
        """
        if not self.base_model:  # a foreign nn.Module defines no `augment`/`embed` contract to honor
            return self.model(im, **kwargs)
        return self.model(im, augment=augment, embed=embed, **kwargs)


class TorchScriptBackend(BaseBackend):
    """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.
    """

    def load_model(self, weight: str | Path) -> None:
        """Load a TorchScript model from a .torchscript file with optional embedded metadata.

        Args:
            weight (str | Path): Path to the .torchscript model file.
        """
        import torchvision  # noqa - required for TorchScript model deserialization

        LOGGER.info(f"Loading {weight} for TorchScript inference...")
        # NNC builds no shape expression for a traced constant, so a repeat forward raises "RuntimeError:
        # _Map_base::at" or segfaults. Never restored: this setter is global, so restoring it races concurrent
        # forwards, and the optimization it disables is the broken one on these versions.
        if TORCH_1_10 and not TORCH_2_1:
            torch._C._jit_set_texpr_fuser_enabled(False)
        self.model = torch.jit.load(weight, map_location=self.device)
        self.model.half() if self.fp16 else self.model.float()
        self.apply_metadata(self.read_metadata(weight))

    def forward(self, im: torch.Tensor) -> torch.Tensor | list[torch.Tensor]:
        """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).
        """
        return self.model(im)
