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

from __future__ import annotations

from pathlib import Path

import torch

from ultralytics.utils import LOGGER
from ultralytics.utils.checks import check_requirements, is_rockchip

from .base import BaseBackend


class RKNNBackend(BaseBackend):
    """Rockchip RKNN inference backend for Rockchip NPU hardware.

    Loads and runs inference with RKNN models (.rknn files) using the RKNN-Toolkit-Lite2 runtime. Only supported on
    Rockchip devices with NPU hardware (e.g., RK3588, RK3566).
    """

    def load_model(self, weight: str | Path) -> None:
        """Load a Rockchip RKNN model from a .rknn file or model directory.

        Args:
            weight (str | Path): Path to the .rknn file or directory containing the model.

        Raises:
            OSError: If not running on a Rockchip device.
            RuntimeError: If model loading or runtime initialization fails.
        """
        if not is_rockchip():
            raise OSError("RKNN inference is only supported on Rockchip devices.")

        LOGGER.info(f"Loading {weight} for RKNN inference...")
        check_requirements("rknn-toolkit-lite2")
        from rknnlite.api import RKNNLite

        w = Path(weight)
        if not w.is_file():
            w = next(w.rglob("*.rknn"))

        self.model = RKNNLite()
        ret = self.model.load_rknn(str(w))
        if ret != 0:
            raise RuntimeError(f"Failed to load RKNN model: {ret}")

        ret = self.model.init_runtime()
        if ret != 0:
            raise RuntimeError(f"Failed to init RKNN runtime: {ret}")

        self.apply_metadata(self.read_metadata(w))

    def forward(self, im: torch.Tensor) -> list:
        """Run inference on the Rockchip NPU.

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

        Returns:
            (list): Model predictions as a list of output arrays.
        """
        h, w = im.shape[1:3]
        im = (im.cpu().numpy() * 255).astype("uint8")
        y = self.model.inference(inputs=[im])
        # INT8 exports use input-relative coordinates so a single per-tensor scale preserves class scores.
        if (
            self.metadata.get("args", {}).get("quantize") == 8
            and self.task in {"detect", "segment", "pose", "obb"}
            and not self.end2end
            and getattr(self, "head", None) != "RTDETRDecoder"
        ):
            kpt_start = 4 + len(self.names)  # pose keypoints follow the box (4) and class-score (nc) channels
            for x in y:
                if x.ndim == 3:
                    x[:, [0, 2]] *= w
                    x[:, [1, 3]] *= h
                    if self.task == "pose":
                        nd = self.kpt_shape[1]  # 2 (x, y) or 3 (x, y, visibility) values per keypoint
                        x[:, kpt_start::nd] *= w
                        x[:, kpt_start + 1 :: nd] *= h
        return y
