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

from __future__ import annotations

import json
import re
import types
from functools import lru_cache
from pathlib import Path

import cv2
import numpy as np
import torch

from ultralytics.utils import ASSETS, IS_JETSON, LOGGER, TORCH_VERSION, ThreadingLocked, imread, is_dgx, is_jetson
from ultralytics.utils.checks import check_requirements, check_tensorrt, check_version
from ultralytics.utils.torch_utils import TORCH_2_4


@lru_cache
def get_tensorrt_logger():
    """Return the shared TensorRT logger, kept alive for every builder and inference runtime."""
    import tensorrt as trt

    return trt.Logger(trt.Logger.INFO)


class _NormalizeCoords(torch.nn.Module):
    """Wrap a model with input-relative box and pose coordinates for per-tensor quantization."""

    def __init__(self, model: torch.nn.Module, h: int, w: int, task: str, nc: int, kpt_shape: tuple | None):
        """Initialize with the wrapped model and prediction metadata."""
        super().__init__()
        self.model = model
        self.h = h
        self.w = w
        self.task = task
        self.nc = nc
        self.kpt_shape = kpt_shape

    def forward(self, x: torch.Tensor):
        """Run the wrapped model and normalize its coordinate channels by input size."""
        y = self.model(x)
        det = y[0] if isinstance(y, (tuple, list)) else y
        box_wh = torch.tensor([self.w, self.h, self.w, self.h], dtype=det.dtype, device=det.device).view(1, 4, 1)
        parts = [det[:, :4] / box_wh]
        if self.task == "pose" and self.kpt_shape:
            parts.append(det[:, 4 : 4 + self.nc])
            b, _, a = det.shape
            kpts = det[:, 4 + self.nc :].view(b, self.kpt_shape[0], self.kpt_shape[1], a)
            kpt_wh = torch.tensor([self.w, self.h], dtype=det.dtype, device=det.device).view(1, 1, 2, 1)
            kpts = torch.cat([kpts[:, :, :2] / kpt_wh, kpts[:, :, 2:]], dim=2)
            parts.append(kpts.reshape(b, -1, a))
        else:
            parts.append(det[:, 4 : 4 + self.nc])
            if det.shape[1] > 4 + self.nc:
                parts.append(det[:, 4 + self.nc :])
        det = torch.cat(parts, dim=1)
        return (det, *y[1:]) if isinstance(y, (tuple, list)) else det


def best_onnx_opset(onnx: types.ModuleType) -> int:
    """Return max ONNX opset for this torch version with ONNX fallback.

    Args:
        onnx (types.ModuleType): The imported `onnx` module, used to cap the opset at the installed ONNX version.

    Returns:
        (int): The ONNX opset version to export with.
    """
    version = ".".join(TORCH_VERSION.split(".")[:2])
    opset = {
        "1.8": 12,
        "1.9": 12,
        "1.10": 13,
        "1.11": 14,
        "1.12": 15,
        "1.13": 17,
        "2.0": 17,  # reduced from 18 to fix ONNX errors
        "2.1": 17,  # reduced from 19
        "2.2": 17,  # reduced from 19
        "2.3": 17,  # reduced from 19
    }.get(version, 18)
    # torch>=2.4 supports opset>=19, but ONNX Runtime CUDA has no Resize-19 or ReduceMax-20 kernel, so opset>=19 runs
    # those nodes on the CPU and copies their tensors back and forth. Its static INT8 quantization also rejects opset>=21.
    return min(opset, onnx.defs.onnx_opset_version())


@ThreadingLocked()
def torch2onnx(
    model: torch.nn.Module,
    im: torch.Tensor | tuple[torch.Tensor, ...],
    output_file: Path | str,
    opset: int = 14,
    input_names: list[str] | None = None,
    output_names: list[str] | None = None,
    dynamic: dict | None = None,
) -> str:
    """Export a PyTorch model to ONNX format.

    Args:
        model (torch.nn.Module): The PyTorch model to export.
        im (torch.Tensor | tuple[torch.Tensor, ...]): Example input tensor(s) for tracing.
        output_file (Path | str): Path to save the exported ONNX file.
        opset (int): ONNX opset version to use for export.
        input_names (list[str] | None): List of input tensor names. Defaults to ``["images"]``.
        output_names (list[str] | None): List of output tensor names. Defaults to ``["output0"]``.
        dynamic (dict | None): Dictionary specifying dynamic axes for inputs and outputs.

    Returns:
        (str): Path to the exported ONNX file.
    """
    if input_names is None:
        input_names = ["images"]
    if output_names is None:
        output_names = ["output0"]
    kwargs = {"dynamo": False} if TORCH_2_4 else {}
    torch.onnx.export(
        model,
        im,
        output_file,
        opset_version=opset,
        input_names=input_names,
        output_names=output_names,
        dynamic_axes=dynamic,
        **kwargs,
    )
    return str(output_file)


def modelopt_quantize_onnx(
    onnx_file: str,
    quantize: int | str | None = None,
    dataset=None,
    shape: tuple[int, int, int, int] = (1, 3, 640, 640),
    dynamic: bool = False,
    prefix: str = "",
) -> str:
    """Bake reduced precision into an ONNX model for TensorRT 11 strongly-typed builds using NVIDIA ModelOpt.

    TensorRT 11 is strongly-typed only: it removed the FP16/INT8 builder flags and the ``IInt8Calibrator`` interface, so
    reduced precision must be expressed in the ONNX graph itself before building. FP16 is applied via ModelOpt AutoCast
    mixed-precision conversion and INT8 via explicit Q/DQ quantization with calibration.

    Args:
        onnx_file (str): Path to the FP32 ONNX file to convert.
        quantize (int | str | None): Precision scheme, 8 for INT8 Q/DQ nodes or 16 for FP16 precision.
        dataset (ultralytics.data.build.InfiniteDataLoader | None): Dataloader providing INT8 calibration images.
            Required when ``quantize=8``.
        shape (tuple[int, int, int, int]): Input shape (batch, channels, height, width) used for INT8 calibration shapes
            of dynamic models and for the FP16 AutoCast calibration image.
        dynamic (bool): Whether the ONNX model uses dynamic input shapes.
        prefix (str): Prefix for log messages.

    Returns:
        (str): Path to the precision-converted ONNX file.

    Raises:
        ValueError: If ``quantize=8`` and no calibration dataset is provided.
    """
    if quantize == 8 and dataset is None:
        raise ValueError("INT8 ModelOpt quantization requires a calibration dataset.")

    # Require modelopt >= 0.44: older releases import onnx.mapping which was removed in onnx >= 1.18 and crash
    check_requirements("nvidia-modelopt[onnx]>=0.44")
    import onnx

    input_name = onnx.load(onnx_file, load_external_data=False).graph.input[0].name
    if quantize == 8:
        from modelopt.onnx.quantization import quantize as modelopt_quantize

        out_file = str(Path(onnx_file).with_suffix(".int8.onnx"))
        # Collect up to ~500 calibration images (TensorRT recommendation); ModelOpt holds them in memory at once,
        # so cap the count to bound memory instead of materializing the entire (possibly thousands-image) dataset.
        images, n = [], 0
        for batch in dataset:
            images.append(batch["img"])
            n += images[-1].shape[0]
            if n >= 512:
                break
        calib = torch.cat(images)
        del images, batch
        calib = calib.to(torch.float32).div_(255.0)
        LOGGER.info(f"{prefix} quantizing ONNX to INT8 with ModelOpt using {calib.shape[0]} calibration images...")
        kwargs = {"calibration_shapes": f"{input_name}:{'x'.join(str(d) for d in shape)}"} if dynamic else {}
        modelopt_quantize(
            onnx_file,
            quantize_mode="int8",
            calibration_data={input_name: calib.cpu().numpy()},
            calibration_method="max",
            # Calibrate on CPU. ModelOpt's CUDA EP session can hit an uncatchable cuDNN-ABI segfault (its pinned
            # onnxruntime-gpu's cuDNN vs the installed torch's) and the TensorRT EP aborts on RTX cards (NvTensorRTRTX);
            # scales are EP-independent, so the INT8 engine is equivalent and only this one-time step is slower.
            calibration_eps=["cpu"],
            # The head's output convolutions, the bare `nn.Conv2d` after each pair of `Conv` blocks, and DFL's fixed
            # conv cost most of the INT8 accuracy for a small share of the runtime, so they stay in float
            nodes_to_exclude=[r".*\.2/Conv$", r".*/dfl/"],
            output_path=out_file,
            **kwargs,
        )
        return out_file

    from modelopt.onnx import autocast

    out_file = str(Path(onnx_file).with_suffix(".fp16.onnx"))
    LOGGER.info(f"{prefix} converting ONNX to FP16 mixed precision with ModelOpt AutoCast...")
    # AutoCast keeps a node in FP32 when its observed activation range exceeds `data_max`, so calibrate it on a real
    # image: unstructured noise inflates the early activations and strands the first convolutions of most models.
    im = cv2.resize(imread(ASSETS / "bus.jpg"), shape[:1:-1])[..., ::-1].transpose(2, 0, 1)  # BGR HWC to RGB CHW
    im = np.resize(im, shape[1:])  # repeat or drop channels for models that are not 3-channel
    im = np.broadcast_to(im, shape).astype(np.float32, order="C") / 255
    onnx.save(
        autocast.convert_to_mixed_precision(
            onnx_file,
            low_precision_type="fp16",
            keep_io_types=True,
            calibration_data={input_name: im},
        ),
        out_file,
    )
    return out_file


def onnx2engine(
    onnx_file: str,
    output_file: Path | str | None = None,
    workspace: float | None = None,
    quantize: int | str | None = None,
    dynamic: bool = False,
    shape: tuple[int, int, int, int] = (1, 3, 640, 640),
    dla: int | None = None,
    dataset=None,
    metadata: dict | None = None,
    verbose: bool = False,
    prefix: str = "",
) -> str:
    """Export a YOLO model to TensorRT engine format.

    Args:
        onnx_file (str): Path to the ONNX file to be converted.
        output_file (Path | str | None): Path to save the generated TensorRT engine file.
        workspace (float | None): Workspace size in GiB for TensorRT, or None for TensorRT auto-allocation.
        quantize (int | str | None): Precision scheme, 16 for FP16 or 8 for INT8.
        dynamic (bool, optional): Enable dynamic input shapes.
        shape (tuple[int, int, int, int], optional): Input shape (batch, channels, height, width).
        dla (int | None): DLA core to use (Jetson devices only).
        dataset (ultralytics.data.build.InfiniteDataLoader, optional): Dataset for INT8 calibration, unused when the
            ONNX graph already carries Q/DQ ranges.
        metadata (dict | None): Metadata to include in the engine file.
        verbose (bool, optional): Enable verbose logging.
        prefix (str, optional): Prefix for log messages.

    Returns:
        (str): Path to the exported engine file.

    Raises:
        ValueError: If INT8 calibration lacks a dataset, or DLA is requested on a non-Jetson device, on TensorRT 11.0,
            or without FP16/INT8 precision.
        RuntimeError: If the ONNX file cannot be parsed or the engine build fails.

    Notes:
        TensorRT version compatibility is handled for workspace size and engine building. On TensorRT 7-10, INT8
        calibration uses an ``IInt8Calibrator`` over ``dataset``, while FP16/INT8 are enabled with builder flags. On
        TensorRT 11 these were removed in favor of strongly-typed networks, so reduced precision is baked into the ONNX
        with NVIDIA ModelOpt before building (FP16 AutoCast, INT8 explicit Q/DQ) by `modelopt_quantize_onnx`. The
        TensorRT 7-10 path keeps the head Sigmoid layers in FP32 to preserve confidence-score calibration (see #24668)
        and the head's output convolutions in FP16 for accuracy. Metadata is serialized and written to the engine file
        if provided.
    """
    import onnx

    # Force re-install TensorRT on CUDA 13 ARM devices to 10.15.x versions for RT-DETR exports
    # https://github.com/ultralytics/ultralytics/issues/22873
    if is_jetson(jetpack=7) or is_dgx():
        check_tensorrt("10.15")

    try:
        import tensorrt as trt
    except ImportError:
        check_tensorrt()
        import tensorrt as trt
    check_version(trt.__version__, ">=7.0.0", hard=True)
    check_version(trt.__version__, "!=10.2.0", msg="https://github.com/ultralytics/ultralytics/pull/24367")

    LOGGER.info(f"\n{prefix} starting export with TensorRT {trt.__version__}...")
    output_file = output_file or Path(onnx_file).with_suffix(".engine")

    logger = get_tensorrt_logger()
    logger.min_severity = trt.Logger.VERBOSE if verbose else trt.Logger.INFO

    # Engine builder
    builder = trt.Builder(logger)
    config = builder.create_builder_config()
    workspace_bytes = int((workspace or 0) * (1 << 30))
    trt_major = int(trt.__version__.split(".", 1)[0])
    is_trt10 = trt_major >= 10
    # TensorRT >= 11 is strongly-typed only: precision builder flags and IInt8Calibrator removed
    is_trt11 = trt_major >= 11
    if workspace_bytes > 0:
        if hasattr(config, "set_memory_pool_limit"):
            config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, workspace_bytes)
        else:  # TensorRT 7 fallback
            config.max_workspace_size = workspace_bytes
    # EXPLICIT_BATCH flag is removed in TensorRT 10 (explicit batch is the only/default mode); keep it for TRT 7/8
    flag = 0 if is_trt10 else (1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH))
    network = builder.create_network(flag)
    # platform_has_fast_fp16/int8 were removed from the Builder in TensorRT 10; default to True when absent
    use_fp16 = getattr(builder, "platform_has_fast_fp16", True) and quantize == 16
    use_int8 = getattr(builder, "platform_has_fast_int8", True) and quantize == 8
    qdq = any(n.op_type == "QuantizeLinear" for n in onnx.load(onnx_file, load_external_data=False).graph.node)
    calibrate = use_int8 and not qdq  # explicit quantization carries its ranges in the graph
    if calibrate and dataset is None:
        raise ValueError("INT8 TensorRT export requires a calibration dataset.")

    # Optionally switch to DLA if enabled
    if dla is not None:
        if not IS_JETSON:
            raise ValueError("DLA is only available on NVIDIA Jetson devices")
        if check_version(trt.__version__, ">=11.0.0,<11.1.0"):
            # DLA is unsupported in TensorRT 11.0 and is planned to return in a later release
            # https://docs.nvidia.com/deeplearning/tensorrt/latest/api/migration/tensorrt-10x-to-11x-jetson.html
            raise ValueError("DLA is not supported in TensorRT 11.0; export with TensorRT 10.x to use DLA.")
        LOGGER.info(f"{prefix} enabling DLA on core {dla}...")
        if not use_fp16 and not use_int8:
            raise ValueError(
                "DLA requires either quantize=16 (FP16) or quantize=8 (INT8). Please enable one of them and try again."
            )
        config.default_device_type = trt.DeviceType.DLA
        config.DLA_core = int(dla)
        config.set_flag(trt.BuilderFlag.GPU_FALLBACK)

    # TensorRT 11 is strongly-typed and removed the FP16/INT8 builder flags and INT8 calibrator, so reduced
    # precision must be baked into the ONNX graph with NVIDIA ModelOpt before parsing (FP16 AutoCast, INT8 Q/DQ)
    if is_trt11 and (use_fp16 or use_int8):
        onnx_file = modelopt_quantize_onnx(onnx_file, 16 if qdq else quantize, dataset, shape, dynamic, prefix)

    # Read ONNX file
    parser = trt.OnnxParser(network, logger)
    if not parser.parse_from_file(onnx_file):
        raise RuntimeError(f"failed to load ONNX file: {onnx_file}")

    # Network inputs
    inputs = [network.get_input(i) for i in range(network.num_inputs)]
    outputs = [network.get_output(i) for i in range(network.num_outputs)]
    for inp in inputs:
        LOGGER.info(f'{prefix} input "{inp.name}" with shape{inp.shape} {inp.dtype}')
    for out in outputs:
        LOGGER.info(f'{prefix} output "{out.name}" with shape{out.shape} {out.dtype}')

    if dynamic:
        profile = builder.create_optimization_profile()
        min_shape = (1, shape[1], 32, 32)  # minimum input shape
        max_shape = (*shape[:2], *(2 * d for d in shape[2:]))  # max input shape, 2x imgsz
        for inp in inputs:
            inp_min = tuple(d if d != -1 else lo for d, lo in zip(inp.shape, min_shape))
            inp_max = tuple(d if d != -1 else hi for d, hi in zip(inp.shape, max_shape))
            profile.set_shape(inp.name, min=inp_min, opt=shape, max=inp_max)
        config.add_optimization_profile(profile)
        if calibrate and not is_trt10:  # deprecated in TensorRT 10, causes internal errors
            config.set_calibration_profile(profile)

    LOGGER.info(
        f"{prefix} building {'INT8' if use_int8 else 'FP' + ('16' if use_fp16 else '32')} engine as {output_file}"
    )
    if use_int8 and not is_trt11:
        config.set_flag(trt.BuilderFlag.INT8)
        config.profiling_verbosity = trt.ProfilingVerbosity.DETAILED
    if (use_fp16 or use_int8 or qdq) and not is_trt11:  # unquantized layers take the fastest precision allowed
        config.set_flag(trt.BuilderFlag.FP16)

    # Explicit Q/DQ graphs need neither calibration nor per-layer Sigmoid constraints.
    if calibrate and not is_trt11:

        class EngineCalibrator(trt.IInt8Calibrator):
            """Custom INT8 calibrator for TensorRT engine optimization.

            This calibrator provides the necessary interface for TensorRT to perform INT8 quantization calibration using
            a dataset. It handles batch generation and calibration algorithm selection.

            Attributes:
                dataset: Dataset for calibration.
                data_iter: Iterator over the calibration dataset.
                algo (trt.CalibrationAlgoType): Calibration algorithm type.
                batch (int): Batch size for calibration.

            Methods:
                get_algorithm: Get the calibration algorithm to use.
                get_batch_size: Get the batch size to use for calibration.
                get_batch: Get the next batch to use for calibration.
                read_calibration_cache: Return no cache so every export calibrates the current model and data.
                write_calibration_cache: Discard the calibration cache.
            """

            def __init__(self, dataset) -> None:  # ultralytics.data.build.InfiniteDataLoader
                """Initialize the INT8 calibrator with a dataset."""
                trt.IInt8Calibrator.__init__(self)
                self.dataset = dataset
                self.data_iter = iter(dataset)
                self.algo = (
                    trt.CalibrationAlgoType.ENTROPY_CALIBRATION_2  # DLA quantization needs ENTROPY_CALIBRATION_2
                    if dla is not None
                    else trt.CalibrationAlgoType.MINMAX_CALIBRATION
                )
                self.batch = dataset.batch_size

            def get_algorithm(self) -> trt.CalibrationAlgoType:
                """Get the calibration algorithm to use."""
                return self.algo

            def get_batch_size(self) -> int:
                """Get the batch size to use for calibration."""
                return self.batch or 1

            def get_batch(self, names) -> list[int] | None:
                """Get the next batch to use for calibration, as a list of device memory pointers."""
                try:
                    im0s = next(self.data_iter)["img"] / 255.0
                    im0s = im0s.to("cuda") if im0s.device.type == "cpu" else im0s
                    return [int(im0s.data_ptr())]
                except StopIteration:
                    # Return None to signal to TensorRT there is no calibration data remaining
                    return None

            def read_calibration_cache(self) -> None:
                """Return no cache so every export calibrates the current model and data."""

            def write_calibration_cache(self, cache: bytes) -> None:
                """Discard the calibration cache, which would be stale for any other model or data."""

        # Load dataset w/ builder (for batching) and calibrate
        config.int8_calibrator = EngineCalibrator(dataset)

        # Implicit quantization cannot exclude op types like ModelOpt on TRT 11, so keep the head Sigmoid (an
        # ACTIVATION layer named after its ONNX node) in FP32 via per-layer precision constraints to preserve
        # confidence-score calibration, mirroring the OpenVINO IgnoredScope
        # https://github.com/ultralytics/ultralytics/issues/24668, and the head's output convolutions and DFL in FP16
        # as `modelopt_quantize_onnx` does. Scope this to the head: every SiLU activation is also a Sigmoid, and
        # constraining all of them costs INT8 speed across backbone and neck.
        names = [network.get_layer(i).name for i in range(network.num_layers)]
        # search, not match: nms=True exports wrap the model in NMSModel, prefixing every name with "/model"
        indices = [int(m.group(1)) for n in names if (m := re.search(r"/model\.(\d+)/", n))]
        head = f"/model.{max(indices)}/" if indices else "/"
        count = 0
        for i in range(network.num_layers):
            layer = network.get_layer(i)
            if head not in layer.name:
                continue
            if layer.type == trt.LayerType.ACTIVATION and "sigmoid" in layer.name.lower():
                dtype = trt.float32
            elif layer.type == trt.LayerType.CONVOLUTION and (layer.name.endswith(".2/Conv") or "/dfl/" in layer.name):
                dtype = trt.float16
            else:
                continue
            layer.precision = dtype
            for j in range(layer.num_outputs):
                layer.set_output_type(j, dtype)
            count += 1
        if count:
            flag = (
                trt.BuilderFlag.OBEY_PRECISION_CONSTRAINTS
                if hasattr(trt.BuilderFlag, "OBEY_PRECISION_CONSTRAINTS")
                else trt.BuilderFlag.STRICT_TYPES
            )
            config.set_flag(flag)  # OBEY_PRECISION_CONSTRAINTS replaced STRICT_TYPES in TensorRT 8.2
            LOGGER.info(f"{prefix} keeping {count} head layers out of INT8 for accuracy")

    # Write file
    if hasattr(builder, "build_serialized_network"):
        engine = builder.build_serialized_network(network, config)
    else:
        engine = builder.build_engine(network, config)
        engine = None if engine is None else engine.serialize()
    if engine is None:
        raise RuntimeError("TensorRT engine build failed, check logs for errors")
    with open(output_file, "wb") as t:
        if metadata is not None:
            meta = json.dumps(metadata)
            t.write(len(meta).to_bytes(4, byteorder="little", signed=True))
            t.write(meta.encode())
        t.write(engine)
    return str(output_file)
