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

from __future__ import annotations

import platform
from copy import deepcopy
from pathlib import Path
from typing import Any

import numpy as np
import torch
from torch import nn

from ultralytics.utils import LINUX, LOGGER, WINDOWS
from ultralytics.utils.checks import check_suffix
from ultralytics.utils.downloads import is_url
from ultralytics.utils.torch_utils import TORCH_1_10, TORCH_1_13, smart_inference_mode

from .backends import (
    AscendBackend,
    AxeleraBackend,
    CoreAIBackend,
    CoreMLBackend,
    DeepXBackend,
    ExecuTorchBackend,
    HailoBackend,
    LiteRTBackend,
    MNNBackend,
    NCNNBackend,
    ONNXBackend,
    ONNXIMXBackend,
    OpenVINOBackend,
    PaddleBackend,
    PyTorchBackend,
    QNNBackend,
    RKNNBackend,
    TensorFlowBackend,
    TensorRTBackend,
    TorchScriptBackend,
    TritonBackend,
)


def check_class_names(names: list | dict) -> dict[int, str]:
    """Check class names and convert to dict format if needed.

    Args:
        names (list | dict): Class names as list or dict format.

    Returns:
        (dict): Class names in dict format with integer keys and string values.

    Raises:
        KeyError: If class indices are invalid for the dataset size.
    """
    if isinstance(names, list):  # names is a list
        names = dict(enumerate(names))  # convert to dict
    if isinstance(names, dict):
        # Convert 1) string keys to int, i.e. '0' to 0, and non-string values to strings, i.e. True to 'True'
        names = {int(k): str(v) for k, v in names.items()}
        n = len(names)
        if not n:
            raise KeyError("0-class dataset, at least one class name is required in your dataset YAML.")
        if max(names.keys()) >= n:
            raise KeyError(
                f"{n}-class dataset requires class indices 0-{n - 1}, but you have invalid class indices "
                f"{min(names.keys())}-{max(names.keys())} defined in your dataset YAML."
            )
        if isinstance(names[0], str) and names[0].startswith("n0"):  # imagenet class codes, i.e. 'n01440764'
            from ultralytics.utils import ROOT, YAML

            names_map = YAML.load(ROOT / "cfg/datasets/ImageNet.yaml")["map"]  # human-readable names
            names = {k: names_map.get(v, v) for k, v in names.items()}
    return names


def default_class_names(data: str | Path | None = None, nc: int = 999) -> dict[int, str]:
    """Load class names from a YAML file or return numerical class names.

    Args:
        data (str | Path, optional): Path to YAML file containing class names.
        nc (int): Number of names to generate when the YAML is missing or unreadable.

    Returns:
        (dict): Dictionary mapping class indices to class names.
    """
    if data:
        try:
            from ultralytics.utils import YAML
            from ultralytics.utils.checks import check_yaml

            return YAML.load(check_yaml(data))["names"]
        except Exception:
            pass
    return {i: f"class{i}" for i in range(nc)}  # return default if above errors


class AutoBackend(nn.Module):
    """Handle dynamic backend selection for running inference using Ultralytics YOLO models.

    The AutoBackend class is designed to provide an abstraction layer for various inference engines. It supports a wide
    range of formats, each with specific naming conventions as outlined below:

        Supported Formats and Naming Conventions:
            | Format                | File Suffix            |
            | --------------------- | ---------------------- |
            | PyTorch               | *.pt                   |
            | TorchScript           | *.torchscript          |
            | ONNX Runtime          | *.onnx                 |
            | ONNX OpenCV DNN       | *.onnx (dnn=True)      |
            | OpenVINO              | *_openvino_model/      |
            | CoreML                | *.mlpackage            |
            | Core AI               | *.aimodel              |
            | TensorRT              | *.engine               |
            | TensorFlow SavedModel | *_saved_model/         |
            | TensorFlow GraphDef   | *.pb                   |
            | TensorFlow Edge TPU   | *_edgetpu.tflite       |
            | LiteRT                | *.tflite               |
            | PaddlePaddle          | *_paddle_model/        |
            | MNN                   | *.mnn                  |
            | NCNN                  | *_ncnn_model/          |
            | IMX                   | *_imx_model/           |
            | RKNN                  | *_rknn_model/          |
            | Triton Inference      | http:// or grpc:// URL |
            | ExecuTorch            | *_executorch_model/    |
            | Axelera AI            | *_axelera_model/       |
            | DEEPX                 | *_deepx_model/         |
            | Qualcomm QNN          | *_qnn.onnx             |
            | Hailo                 | *_hailo_model/         |
            | Huawei Ascend         | *_ascend_model/        |
            | AMD Xilinx            | *_xilinx_model/        |

    Attributes:
        backend (BaseBackend): The loaded inference backend instance.
        format (str): The model format (e.g., 'pt', 'onnx', 'engine').
        model: The underlying model, delegated from `backend.model` (nn.Module for PyTorch, runtime object otherwise).
        device (torch.device): The device (CPU or GPU) on which the model is loaded.
        task (str): The type of task the model performs (detect, segment, semantic, depth, classify, pose, obb).
        names (dict): A dictionary of class names that the model can detect.
        stride (int): The model stride, typically 32 for YOLO models.
        fp16 (bool): Whether the model uses half-precision (FP16) inference.
        nhwc (bool): Whether the model expects NHWC input format instead of NCHW.

    Methods:
        forward: Run inference on an input image.
        from_numpy: Convert NumPy arrays to tensors on the model device.
        warmup: Warm up the model with a dummy input.
        _model_type: Determine the model type from file path.

    Examples:
        >>> import torch
        >>> model = AutoBackend(model="yolo26n.pt", device=torch.device("cpu"))
        >>> preds = model(torch.zeros(1, 3, 640, 640))
    """

    _BACKEND_MAP = {
        "pt": PyTorchBackend,
        "torchscript": TorchScriptBackend,
        "onnx": ONNXBackend,
        "dnn": ONNXBackend,  # Special case: ONNX with DNN
        "openvino": OpenVINOBackend,
        "engine": TensorRTBackend,
        "coreml": CoreMLBackend,
        "coreai": CoreAIBackend,
        "saved_model": TensorFlowBackend,
        "pb": TensorFlowBackend,
        "edgetpu": TensorFlowBackend,
        "litert": LiteRTBackend,
        "paddle": PaddleBackend,
        "mnn": MNNBackend,
        "ncnn": NCNNBackend,
        "imx": ONNXIMXBackend,
        "rknn": RKNNBackend,
        "triton": TritonBackend,
        "executorch": ExecuTorchBackend,
        "axelera": AxeleraBackend,
        "deepx": DeepXBackend,
        "qnn": QNNBackend,
        "hailo": HailoBackend,
        "ascend": AscendBackend,
        "xilinx": ONNXBackend,
    }

    @smart_inference_mode(False)
    def __init__(
        self,
        model: str | Path | torch.nn.Module = "yolo26n.pt",
        device: torch.device | str | None = None,
        dnn: bool = False,
        data: str | Path | None = None,
        fp16: bool = False,
        fuse: bool = True,
        verbose: bool = True,
        channels_last: bool | None = None,
        end2end: bool | None = None,
    ):
        """Initialize the AutoBackend for inference.

        Args:
            model (str | Path | torch.nn.Module): Path to the model weights file or a module instance.
            device (torch.device | str, optional): Device to run the model on, or a 'tpu', 'intel' or 'vulkan' device
                string from `select_device`. Defaults to CPU when None.
            dnn (bool): Use OpenCV DNN module for ONNX inference.
            data (str | Path, optional): Path to the additional data.yaml file containing class names.
            fp16 (bool): Enable half-precision inference. Supported only on specific backends.
            fuse (bool): Fuse Conv2D + BatchNorm layers for optimization.
            verbose (bool): Enable verbose logging.
            channels_last (bool, optional): Use channels-last memory format, or auto-enable it on supported x86 CPUs.
            end2end (bool, optional): Select the native detection head before fusion; None preserves its current mode.
        """
        super().__init__()
        device = device or torch.device("cpu")
        # Determine model format from path/URL
        format = "pt" if isinstance(model, nn.Module) else self._model_type(model, dnn)
        if (
            isinstance(model, nn.Module)
            and TORCH_1_10
            and any(x.is_inference() for x in (*model.parameters(), *model.buffers()))
        ):
            model = deepcopy(model)  # retained backends require normal tensors for fusion and later mutation

        # Check if format supports FP16
        fp16 &= format in {"pt", "torchscript", "onnx", "openvino", "engine"}

        # Set device
        if (
            isinstance(device, torch.device)
            and torch.cuda.is_available()
            and device.type != "cpu"
            and format not in {"pt", "torchscript", "engine", "onnx", "paddle"}
        ):
            device = torch.device("cpu")

        # Select and initialize the appropriate backend
        backend_kwargs = {"device": device, "fp16": fp16}

        if format not in self._BACKEND_MAP:
            from ultralytics.engine.exporter import export_formats

            raise TypeError(
                f"model='{model}' is not a supported model format. "
                f"Ultralytics supports: {export_formats()['Format']}\n"
                f"See https://docs.ultralytics.com/modes/predict for help."
            )
        if format == "pt":
            backend_kwargs["fuse"] = fuse
            backend_kwargs["verbose"] = verbose
            backend_kwargs["end2end"] = end2end
        elif format in {"saved_model", "pb", "edgetpu", "dnn"}:
            backend_kwargs["format"] = format
        self.backend = self._BACKEND_MAP[format](model, **backend_kwargs)

        if format == "pt":
            device_type = torch.device(self.backend.device).type
            supported = device_type == "cuda" or (
                TORCH_1_13
                and device_type == "cpu"
                and platform.machine() in {"AMD64", "x86_64"}
                and torch.backends.mkldnn.is_available()
                and torch.backends.mkldnn.enabled
            )
            if channels_last is None:
                channels_last = device_type == "cpu" and supported and (LINUX or WINDOWS)
            if channels_last and not supported:
                LOGGER.warning(f"'channels_last=True' is not supported on '{device_type}', ignoring.")
            self.backend.model.to(
                memory_format=torch.channels_last if channels_last and supported else torch.contiguous_format
            )
        elif channels_last:
            LOGGER.warning(f"'channels_last=True' applies only to native PyTorch models, ignoring format='{format}'.")

        self.nhwc = format in {"coreml", "saved_model", "pb", "edgetpu", "rknn"}
        self.format = format

        # Ensure backend has names (fallback to default if not set by metadata)
        if not self.backend.names:
            self.backend.names = default_class_names(data)
        self.backend.names = check_class_names(self.backend.names)
        empty = [k for k, v in self.backend.names.items() if not v.strip()]
        if empty:
            LOGGER.warning(f"Empty class name string(s) at class indices {empty} will display as blank labels.")

    def __getattr__(self, name: str) -> Any:
        """Delegate attribute access to the backend.

        This allows AutoBackend to transparently expose backend attributes
        without explicit copying.

        Args:
            name (str): Attribute name to look up.

        Returns:
            (Any): The attribute value from the backend.

        Raises:
            AttributeError: If the attribute is not found in backend.
        """
        if "backend" in self.__dict__ and hasattr(self.backend, name):
            return getattr(self.backend, name)
        return super().__getattr__(name)

    def forward(
        self,
        im: torch.Tensor,
        augment: bool = False,
        embed: list | None = None,
        **kwargs: Any,
    ) -> Any:
        """Run inference on an AutoBackend model.

        Args:
            im (torch.Tensor): The image tensor to perform inference on.
            augment (bool): Whether to apply test-time augmentation (native PyTorch models only).
            embed (list, optional): A list of layer indices to return embeddings from (native PyTorch models only).
            **kwargs (Any): Additional keyword arguments passed to native PyTorch models; ignored by other formats.

        Returns:
            (Any): The raw model output, with NumPy arrays converted to tensors on `self.device`.
        """
        if self.nhwc:
            im = im.permute(0, 2, 3, 1)  # torch BCHW to numpy BHWC shape(1,320,192,3)
        if self.backend.fp16 and im.dtype != torch.float16:
            im = im.half()
        fixed = not self.metadata.get("dynamic") and self.format not in {"torchscript", "ncnn", "deepx", "axelera"}
        if (pad := self.batch - im.shape[0] if fixed else 0) > 0:  # static-batch exports reject short batches
            im = torch.cat((im, im.new_zeros(pad, *im.shape[1:])))

        # Build forward kwargs based on backend type
        forward_kwargs = {}
        if self.format == "pt":
            forward_kwargs = {"augment": augment, "embed": embed, **kwargs}

        y = self.backend.forward(im, **forward_kwargs)
        if pad > 0:  # drop the zero-padded rows
            y = [x[:-pad] for x in y] if isinstance(y, (list, tuple)) else y[:-pad]

        if isinstance(y, (list, tuple)):
            if len(self.names) == 999 and (self.task == "segment" or len(y) == 2):  # segments and names not defined
                nc = y[0].shape[1] - y[1].shape[1] - 4  # y = (1, 116, 8400), (1, 32, 160, 160)
                self.names = {i: f"class{i}" for i in range(nc)}
            return self.from_numpy(y[0]) if len(y) == 1 else [self.from_numpy(x) for x in y]
        else:
            return self.from_numpy(y)

    def from_numpy(self, x: Any) -> Any:
        """Normalize a backend output to the model device when possible.

        Args:
            x (Any): Backend output to normalize.

        Returns:
            (Any): Tensor on `self.device`, or the unchanged non-tensor output.
        """
        if isinstance(x, np.ndarray):
            return torch.as_tensor(x, device=self.device)  # shares memory on CPU, one fused copy to accelerators
        return x.to(self.device) if isinstance(x, torch.Tensor) else x

    def warmup(self, imgsz: tuple[int, int, int, int] = (1, 3, 640, 640), im: torch.Tensor | None = None) -> None:
        """Warm up the model by running forward pass(es).

        Args:
            imgsz (tuple[int, int, int, int]): Dummy input shape in (batch, channels, height, width) format.
            im (torch.Tensor, optional): Input tensor to reuse instead of allocating a dummy.
        """
        from ultralytics.utils.nms import non_max_suppression

        if not self.end2end:
            import torchvision  # noqa (import here triggers torchvision NMS use in nms.py)
        if self.format in {"pt", "torchscript", "onnx", "engine", "saved_model", "pb", "triton"} and (
            self.device.type != "cpu" or self.format == "triton"
        ):
            im = (
                im
                if im is not None
                else torch.empty(*imgsz, dtype=torch.half if self.fp16 else torch.float, device=self.device)
            )
            for _ in range(2 if self.format == "torchscript" else 1):
                self.forward(im)  # warmup model
                warmup_boxes = torch.rand(1, 84, 16, device=self.device)  # 16 boxes works best empirically
                warmup_boxes[:, :4] *= im.shape[-1]
                non_max_suppression(warmup_boxes)  # warmup NMS

    @staticmethod
    def _model_type(p: str | Path = "path/to/model.pt", dnn: bool = False) -> str:
        """Take a path to a model file and return the model format string.

        Args:
            p (str | Path): Path to the model file or Triton URL.
            dnn (bool): Whether to use OpenCV DNN module for ONNX inference.

        Returns:
            (str): Model format string (e.g., 'pt', 'onnx', 'engine', 'triton').

        Examples:
            >>> fmt = AutoBackend._model_type("path/to/model.onnx")
            >>> assert fmt == "onnx"
        """
        from ultralytics.engine.exporter import export_formats

        sf = export_formats()["Suffix"]
        if not is_url(p) and not isinstance(p, str):
            check_suffix(p, sf)
        name = Path(p).name
        # The suffix ending last wins, then the longest, i.e. 'best.pt.onnx' -> onnx, 'best_qnn.onnx' -> qnn
        matches = [
            (name.rfind(s) + len(s), len(s), f)
            for s, f in zip([*sf, ".mlmodel"], [*export_formats()["Argument"], "coreml"])
            if s in name
        ]
        format = max(matches)[2] if matches else None
        if format == "-":
            format = "pt"
        elif format == "onnx" and dnn:
            format = "dnn"
        elif format is None:
            from urllib.parse import urlsplit

            url = urlsplit(p)
            if bool(url.netloc) and bool(url.path) and url.scheme in {"http", "grpc"}:
                format = "triton"
        return format

    def eval(self) -> AutoBackend:
        """Set the backend model to evaluation mode if supported."""
        if hasattr(self.backend, "model") and hasattr(self.backend.model, "eval"):
            self.backend.model.eval()
        return super().eval()

    def _apply(self, fn) -> AutoBackend:
        """Apply a function to backend.model parameters, buffers, and tensors.

        This method extends the functionality of the parent class's _apply method by additionally applying the
        function to the backend model and updating the backend device. It's typically used for operations like moving
        the model to a different device or changing its precision.

        Args:
            fn (Callable): A function to be applied to the model's tensors. This is typically a method like to(), cpu(),
                cuda(), half(), or float().

        Returns:
            (AutoBackend): The model instance with the function applied and updated attributes.
        """
        super()._apply(fn)
        if hasattr(self.backend, "model") and isinstance(self.backend.model, nn.Module):
            self.backend.model._apply(fn)
            self.backend.device = next(self.backend.model.parameters()).device  # update device after move
        return self
