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

from __future__ import annotations

import functools
import gc
import math
import os
import random
import time
from contextlib import contextmanager
from copy import deepcopy
from datetime import datetime
from pathlib import Path
from typing import Any

import numpy as np
import torch
import torch.distributed as dist
import torch.nn.functional as F
from torch import nn
from torch.nn.utils.fusion import fuse_conv_bn_weights
from torch.utils._pytree import tree_flatten

from ultralytics import __version__
from ultralytics.utils import (
    ARM64,
    DEFAULT_CFG_DICT,
    DEFAULT_CFG_KEYS,
    LOCAL_RANK,
    LOGGER,
    MACOS,
    NUM_THREADS,
    PYTHON_VERSION,
    TORCH_VERSION,
    TORCHVISION_VERSION,
    WINDOWS,
    colorstr,
)
from ultralytics.utils.checks import check_requirements, check_version
from ultralytics.utils.cpu import CPUInfo
from ultralytics.utils.patches import torch_load

# Version checks (all default to version>=min_version)
TORCH_1_9 = check_version(TORCH_VERSION, "1.9.0")
TORCH_1_10 = check_version(TORCH_VERSION, "1.10.0")
TORCH_1_11 = check_version(TORCH_VERSION, "1.11.0")
TORCH_1_13 = check_version(TORCH_VERSION, "1.13.0")
TORCH_2_0 = check_version(TORCH_VERSION, "2.0.0")
TORCH_2_1 = check_version(TORCH_VERSION, "2.1.0")
TORCH_2_2 = check_version(TORCH_VERSION, "2.2.0")
TORCH_2_3 = check_version(TORCH_VERSION, "2.3.0")
TORCH_2_4 = check_version(TORCH_VERSION, "2.4.0")
TORCH_2_5 = check_version(TORCH_VERSION, "2.5.0")
TORCH_2_7 = check_version(TORCH_VERSION, "2.7.0")
TORCH_2_8 = check_version(TORCH_VERSION, "2.8.0")
TORCH_2_9 = check_version(TORCH_VERSION, "2.9.0")
TORCH_2_10 = check_version(TORCH_VERSION, "2.10.0")
TORCH_2_12 = check_version(TORCH_VERSION, "2.12.0")
TORCH_2_13 = check_version(TORCH_VERSION, "2.13.0")
TORCHVISION_0_10 = check_version(TORCHVISION_VERSION, "0.10.0")
TORCHVISION_0_11 = check_version(TORCHVISION_VERSION, "0.11.0")
TORCHVISION_0_13 = check_version(TORCHVISION_VERSION, "0.13.0")
TORCHVISION_0_18 = check_version(TORCHVISION_VERSION, "0.18.0")
if WINDOWS and check_version(TORCH_VERSION, "==2.4.0"):  # reject version 2.4.0 on Windows
    LOGGER.warning(
        "Known issue with torch==2.4.0 on Windows with CPU, recommend upgrading to torch>=2.4.1 to resolve "
        "https://github.com/ultralytics/ultralytics/issues/15049"
    )


def get_torch_device_backend(device: torch.device | str):
    """Return the PyTorch module that owns the selected device backend."""
    device_type = getattr(device, "type", str(device).split(":")[0])
    return torch.get_device_module(device_type) if hasattr(torch, "get_device_module") else getattr(torch, device_type)


@contextmanager
def torch_distributed_zero_first(local_rank: int):
    """Ensure all processes in distributed training wait for the local master (rank 0) to complete a task first."""
    initialized = dist.is_available() and dist.is_initialized()
    use_ids = initialized and dist.get_backend() == "nccl"

    if initialized and local_rank not in {-1, 0}:
        dist.barrier(device_ids=[torch.cuda.current_device()]) if use_ids else dist.barrier()
    yield
    if initialized and local_rank == 0:
        dist.barrier(device_ids=[torch.cuda.current_device()]) if use_ids else dist.barrier()


def smart_inference_mode(mode=True):
    """Apply or disable torch inference mode while supporting the minimum torch version."""

    def decorate(fn):
        """Apply appropriate torch decorator for inference mode based on torch version."""
        if not mode:
            if TORCH_1_9:

                @functools.wraps(fn)
                def disable(*args, **kwargs):
                    with torch.inference_mode(False), torch.no_grad():
                        return fn(*args, **kwargs)

                return disable
            return torch.no_grad()(fn)
        if TORCH_1_9 and torch.is_inference_mode_enabled():
            return fn  # already in inference_mode, act as a pass-through
        else:
            return (torch.inference_mode if TORCH_1_10 else torch.no_grad)()(fn)

    return decorate


def autocast(enabled: bool | torch.dtype, device: str = "cuda"):
    """Get the appropriate autocast context manager based on PyTorch version and AMP setting.

    This function returns a context manager for automatic mixed precision (AMP) training that is compatible with both
    older and newer versions of PyTorch. It handles the differences in the autocast API between PyTorch versions.

    Args:
        enabled (bool | torch.dtype): Whether to enable AMP, or the autocast dtype to enable.
        device (str, optional): Device type to use for autocast, e.g. "cuda" or "npu".

    Returns:
        (torch.amp.autocast): The appropriate autocast context manager.

    Raises:
        RuntimeError: If bfloat16 is requested without torch>=1.13 and a CUDA device with native bfloat16 support.

    Examples:
        >>> from ultralytics.utils.torch_utils import autocast
        >>> with autocast(enabled=True):
        ...     # Your mixed precision operations here
        ...     pass

    Notes:
        Uses `torch.amp.autocast` on torch>=1.13 and the backend-specific AMP context on older releases.
    """
    dtype = enabled if isinstance(enabled, torch.dtype) else None
    enabled = bool(enabled)
    if dtype is torch.bfloat16:
        bf16_supported = device == "cuda" and TORCH_1_13
        if bf16_supported:
            bf16_supported = (
                torch.cuda.is_bf16_supported(including_emulation=False)
                if TORCH_2_4
                else torch.cuda.is_bf16_supported()
                and (bool(torch.version.hip) or torch.cuda.get_device_capability()[0] >= 8)
            )
        if not bf16_supported:
            raise RuntimeError("bfloat16 autocast requires CUDA with native bfloat16 support and torch>=1.13")
    kwargs = {"dtype": dtype} if dtype is not None else {}
    if device == "npu":
        import torch_npu

        return torch_npu.npu.amp.autocast(enabled=enabled, **kwargs)
    if TORCH_1_13:
        if device == "mps" and not TORCH_2_5:  # MPS autocast added in torch 2.5.0, errors on older versions
            device, enabled = "cpu", False
        return torch.amp.autocast(device, enabled=enabled, **kwargs)
    else:
        return torch.cuda.amp.autocast(enabled and device == "cuda")


@functools.lru_cache
def get_cpu_info():
    """Return a string with system CPU information, e.g. 'Apple M2'."""
    return CPUInfo.name()


@functools.lru_cache
def get_gpu_info(index):
    """Return a string with system GPU information, e.g. 'Tesla T4, 15102MiB'."""
    properties = torch.cuda.get_device_properties(index)
    return f"{properties.name}, {properties.total_memory / (1 << 20):.0f}MiB"


def parse_device(device: str | int | list | tuple | torch.device = "") -> str:
    """Parse a device request of any form into a canonical device string.

    Args:
        device (str | int | list | tuple | torch.device, optional): Device request, e.g. 'cuda:0', '0,1', [0, 1], 'cpu',
            'mps', or '-1' to auto-select an idle GPU ('-1,-1' for two).

    Returns:
        (str): Canonical device string, e.g. '', 'cpu', 'mps', '0', or '0,1'.

    Examples:
        >>> parse_device("cuda:0")
        '0'

        >>> parse_device([0, 1])
        '0,1'

    Notes:
        Each '-1' is replaced with an idle GPU index. Requested ids exceeding the torch device count that match
        physical GPU ids visible under an external CUDA_VISIBLE_DEVICES restriction are translated to the
        corresponding torch indices, e.g. '3' -> '0' when CUDA_VISIBLE_DEVICES='3'; in-range ids are always torch
        indices, keeping parsing idempotent. Returned indices are relative to the active restriction, so strings
        persisted under one environment (e.g. resumed checkpoint args) address the same physical GPUs only in that
        environment.
    """
    if isinstance(device, torch.device):
        if device.type == "cuda" and device.index is None:
            return ""  # indexless torch.device('cuda') means the current CUDA device, i.e. the '' default request
        if device.type == "cpu":
            return "cpu"  # an indexed torch.device('cpu', 0) is the same cpu
        if device.type in {"npu", "xpu"}:
            return device.type if device.index is None else f"{device.type}:{device.index}"
    device = str(device).lower()
    for remove in "cuda:", "none", "(", ")", "[", "]", "'", " ":
        device = device.replace(remove, "")  # to string, 'cuda:0' -> '0' and '(0, 1)' -> '0,1'
    if device == "cuda":
        device = "0"
    for backend in ("npu", "xpu"):
        if device.startswith(backend):
            indices = device[len(backend) :].lstrip(":").replace(f"{backend}:", "")
            indices = ",".join(str(int(x)) if x.isdigit() else x for x in indices.split(",") if x)
            return f"{backend}:{indices}" if indices else backend
    device = ",".join(str(int(x)) if x.isdigit() else x for x in device.split(",") if x)  # "0,,01" -> "0,1"
    # Visible physical ids normalized like requested ids and truncated to the torch device count, mirroring CUDA's
    # atoi-style parsing and its stop at the first invalid CVD entry
    cvd = os.environ.get("CUDA_VISIBLE_DEVICES", "").replace(" ", "")
    visible = [str(int(x)) if x.isdigit() else x for x in cvd.split(",") if x][: torch.cuda.device_count()]
    indices = [x for x in device.split(",") if x.isdigit()]  # requested ids, excluding '-1' and non-numeric tokens
    if indices and all(x in visible for x in indices) and any(int(x) >= torch.cuda.device_count() for x in indices):
        # Ids exceeding the torch device count can only be physical GPU ids under an external CUDA_VISIBLE_DEVICES
        # restriction -> translate to torch indices; in-range ids are torch indices, keeping repeated parses stable
        device = ",".join(str(visible.index(x)) if x.isdigit() else x for x in device.split(","))
    if "-1" in device:
        from ultralytics.utils.autodevice import GPUInfo

        # Replace each -1 with an idle GPU or remove it; GPUInfo searches physical NVML ids among externally visible
        # GPUs only, translated back to torch indices under a CUDA_VISIBLE_DEVICES restriction
        parts = device.split(",")
        candidates = [int(x) for x in visible if x.isdigit()] if visible else None
        selected = GPUInfo().select_idle_gpu(count=parts.count("-1"), min_memory_fraction=0.2, indices=candidates)
        selected = [visible.index(str(x)) for x in selected] if visible else selected
        for i in range(len(parts)):
            if parts[i] == "-1":
                parts[i] = str(selected.pop(0)) if selected else ""
        device = ",".join(p for p in parts if p)
    return device


def select_device(device="", newline=False, verbose=True):
    """Select the appropriate PyTorch device based on the provided arguments.

    The function takes a device request (see `parse_device`) or a torch.device object and returns a torch.device object
    representing the selected device. The function also validates the number of available devices and raises an
    exception if the requested device(s) are not available.

    Args:
        device (str | int | list | tuple | torch.device, optional): Device request or torch.device object. Options
            include 'cpu', 'cuda', '0', '0,1,2,3', 'mps', 'npu:0', 'npu:0,1', 'xpu:0', 'xpu:0,1', or '-1' to auto-select
            an idle GPU. Defaults to the current CUDA device, or CPU if no GPU is available.
        newline (bool, optional): If True, adds a newline at the end of the log string.
        verbose (bool, optional): If True, logs the device information.

    Returns:
        (torch.device | str): Selected device. For multi-GPU requests this is the first GPU. 'tpu', 'intel', and
            'vulkan' device strings and other non-cpu/cuda/npu/xpu torch.device inputs are returned unchanged.

    Raises:
        ValueError: If the requested CUDA, NPU, or XPU device(s) are invalid or unavailable.

    Examples:
        >>> select_device("cuda:0")
        device(type='cuda', index=0)

        >>> select_device("cpu")
        device(type='cpu')

    Notes:
        CUDA indices are torch device indices, which reflect any externally set CUDA_VISIBLE_DEVICES. This function
        never modifies CUDA_VISIBLE_DEVICES; an explicit single-GPU request is made the default CUDA device with
        torch.cuda.set_device() so that indexless 'cuda' operations land on it, while default '' requests (resolved
        to the current device) and multi-GPU requests (DDP ranks pin their own device in trainer._setup_ddp()) leave
        the current device untouched.
    """
    if isinstance(device, torch.device):
        if device.type not in {"cpu", "cuda", "npu", "xpu"}:
            return device  # other torch.device inputs pass through; cpu and accelerator inputs canonicalize below
    elif str(device).startswith(("tpu", "intel", "vulkan")):
        return device

    s = f"Ultralytics {__version__} 🚀 Python-{PYTHON_VERSION} torch-{TORCH_VERSION} "
    device = parse_device(device)

    if device.startswith(("npu", "xpu")):
        device_type = device.split(":", 1)[0]
        if device_type == "npu":
            try:
                import torch_npu  # noqa
            except ImportError:
                raise ValueError(
                    f"Invalid NPU 'device={device}'. Install 'torch_npu' at https://github.com/Ascend/pytorch"
                )
        if not hasattr(torch, device_type):
            raise ValueError(f"Invalid {device_type.upper()} 'device={device}' requested. Backend is not available.")
        backend = get_torch_device_backend(device_type)
        if not backend.is_available():
            raise ValueError(f"Invalid {device_type.upper()} 'device={device}' requested. Backend is not available.")

        requested = ["0"] if device == device_type else device[4:].split(",")
        indices = [int(x) for x in requested if x.isdigit()]
        if not indices or len(indices) != len(requested) or len(indices) != len(set(indices)):
            raise ValueError(
                f"Invalid {device_type.upper()} 'device={device}' format. "
                f"Use '{device_type}', '{device_type}:0', or '{device_type}:0,1'."
            )
        n = backend.device_count()
        if any(idx >= n for idx in indices):
            raise ValueError(
                f"Invalid {device_type.upper()} 'device={device}' requested. Only {n} device(s) available."
            )

        if len(indices) == 1:
            backend.set_device(indices[0])  # multi-device DDP ranks each pin their device in trainer._setup_ddp()
        if verbose:
            space = " " * len(s)
            for i, idx in enumerate(indices):
                s += f"{'' if i == 0 else space}{device_type.upper()}:{idx} ({backend.get_device_name(idx)})\n"
            LOGGER.info(s if newline else s.rstrip())
        return torch.device(device_type, indices[0])

    cpu = device == "cpu"
    mps = device in {"mps", "mps:0"}  # Apple Metal Performance Shaders (MPS)
    if not cpu and not mps and device:  # non-cpu device requested
        valid = all(x.isdigit() and int(x) < torch.cuda.device_count() for x in device.split(","))
        if not (torch.cuda.is_available() and valid):
            LOGGER.info(s)
            install = (
                "See https://pytorch.org/get-started/locally/ for up-to-date torch install instructions if no "
                "CUDA devices are seen by torch.\n"
                if torch.cuda.device_count() == 0
                else ""
            )
            raise ValueError(
                f"Invalid CUDA 'device={device}' requested."
                f" Use 'device=cpu' or pass valid CUDA device(s) if available,"
                f" i.e. 'device=0' or 'device=0,1,2,3' for Multi-GPU.\n"
                f"\ntorch.cuda.is_available(): {torch.cuda.is_available()}"
                f"\ntorch.cuda.device_count(): {torch.cuda.device_count()}"
                f"\nos.environ['CUDA_VISIBLE_DEVICES']: {os.environ.get('CUDA_VISIBLE_DEVICES')}\n"
                f"{install}"
            )

    if not cpu and not mps and torch.cuda.is_available():  # prefer GPU if available
        devices = device.split(",") if device else [str(torch.cuda.current_device())]  # '' -> current default device
        space = " " * len(s)
        for i, d in enumerate(devices):
            s += f"{'' if i == 0 else space}CUDA:{d} ({get_gpu_info(int(d))})\n"
        arg = f"cuda:{devices[0]}"
        if device and len(devices) == 1:  # explicit single-GPU request only: '' never moves the current device, and
            torch.cuda.set_device(int(devices[0]))  # multi-GPU DDP ranks each pin their own device in _setup_ddp()
    elif mps and TORCH_2_0 and torch.backends.mps.is_available():
        # Prefer MPS if available
        s += f"MPS ({get_cpu_info()})\n"
        arg = "mps"
    else:  # revert to CPU
        s += f"CPU ({get_cpu_info()})\n"
        arg = "cpu"

    if arg in {"cpu", "mps"}:
        torch.set_num_threads(NUM_THREADS)  # reset OMP_NUM_THREADS for cpu training
    if arg == "cpu" and MACOS and ARM64 and TORCH_2_3:
        torch.backends.nnpack.set_flags(False)  # NNPACK conv2d at batch>=16 is 6x slower than im2col on Apple silicon
    if verbose:
        LOGGER.info(s if newline else s.rstrip())
    return torch.device(arg)


def time_sync(device: torch.device | None = None):
    """Return PyTorch-accurate time, synchronizing the accelerator first unless the device is CPU or MPS."""
    if device is None or device.type not in {"cpu", "mps"}:
        accelerator = get_torch_device_backend(device or "cuda")
        if accelerator.is_available() and hasattr(accelerator, "synchronize"):
            accelerator.synchronize()
    return time.perf_counter()


def fuse_conv_and_bn(conv, bn):
    """Fuse Conv2d and BatchNorm2d layers for inference optimization.

    Args:
        conv (nn.Conv2d): Convolutional layer to fuse.
        bn (nn.BatchNorm2d): Batch normalization layer to fuse.

    Returns:
        (nn.Conv2d): The fused convolutional layer with gradients disabled.

    Examples:
        >>> conv = nn.Conv2d(3, 16, 3)
        >>> bn = nn.BatchNorm2d(16)
        >>> fused_conv = fuse_conv_and_bn(conv, bn)
    """
    conv.weight, conv.bias = fuse_conv_bn_weights(
        conv.weight, conv.bias, bn.running_mean, bn.running_var, bn.eps, bn.weight, bn.bias
    )
    q = getattr(conv, "weight_quantizer", None)
    if q is not None:  # QAT: the per-channel weight range scales with the folded BN, so the INT8 codes do not move
        q.amax = q.amax * (bn.weight / torch.sqrt(bn.running_var + bn.eps)).abs().view_as(q.amax)
    return conv.requires_grad_(False)


def fuse_deconv_and_bn(deconv, bn):
    """Fuse ConvTranspose2d and BatchNorm2d layers for inference optimization.

    Args:
        deconv (nn.ConvTranspose2d): Transposed convolutional layer to fuse.
        bn (nn.BatchNorm2d): Batch normalization layer to fuse.

    Returns:
        (nn.ConvTranspose2d): The fused transposed convolutional layer with gradients disabled.

    Examples:
        >>> deconv = nn.ConvTranspose2d(16, 3, 3)
        >>> bn = nn.BatchNorm2d(3)
        >>> fused_deconv = fuse_deconv_and_bn(deconv, bn)
    """
    if isinstance(bn, nn.Identity):  # ConvTranspose(bn=False) leaves bn as nn.Identity, nothing to fuse
        return deconv.requires_grad_(False)
    # ConvTranspose2d weight is [in_channels, out_channels // groups, kH, kW]; view it in the Conv2d layout
    # [out_channels, in_channels // groups, kH, kW] so the per-output-channel BN scale folds along axis 0
    g, (ci, co, *k) = deconv.groups, deconv.weight.shape
    weight = deconv.weight.view(g, ci // g, co, *k).transpose(1, 2).reshape(g * co, ci // g, *k)
    weight, deconv.bias = fuse_conv_bn_weights(
        weight, deconv.bias, bn.running_mean, bn.running_var, bn.eps, bn.weight, bn.bias
    )
    deconv.weight = nn.Parameter(weight.view(g, co, ci // g, *k).transpose(1, 2).reshape(ci, co, *k))
    return deconv.requires_grad_(False)


# ModelOpt's torch plugins import huggingface_hub unconditionally but declare it only under its heavy [hf] extra,
# so a bare install cannot import modelopt.torch at all
MODELOPT_REQUIREMENTS = ["nvidia-modelopt>=0.44", "huggingface_hub"]


def prepare_qat(model: nn.Module, dataloader, preprocess, batches: int = 8) -> nn.Module:
    """Insert INT8 fake-quantization into a model for quantization-aware training (QAT).

    Swaps Conv and Linear layers for ModelOpt equivalents that fake-quantize their input and weight, so training learns
    weights that survive INT8 export and `torch.onnx.export` emits those ranges as Q/DQ nodes. Activation and weight
    ranges are calibrated once from `batches` batches and then held fixed (ModelOpt's INT8 config keeps `amax` as a
    buffer, not a learnable parameter), so training adapts the weights to them.

    Training keeps BatchNorm unfused; `fuse()` folds it at export and rescales the weight ranges along. The head's
    output layers, the bare convolutions and linears outside its `Conv` blocks, are left in float to limit INT8 accuracy
    loss.

    Args:
        model (nn.Module): Model to prepare, modified in place.
        dataloader (Iterable): Loader yielding Ultralytics batches for the initial range calibration.
        preprocess (Callable): Task trainer preprocessing applied to each calibration batch.
        batches (int): Number of calibration batches.

    Returns:
        (nn.Module): The prepared model, carrying fake-quantization modules.
    """
    with torch_distributed_zero_first(LOCAL_RANK):
        check_requirements(MODELOPT_REQUIREMENTS)
        import modelopt.torch.quantization as mtq

    def forward_loop(m):
        """Calibrate through the task batch path."""
        for batch, _ in zip(dataloader, range(batches)):
            m(preprocess(batch))

    LOGGER.info(f"Preparing INT8 quantization-aware training from {batches} calibration batches...")
    training = model.training
    model.eval()  # freeze BatchNorm statistics
    with torch.no_grad():
        model = mtq.quantize(model, mtq.INT8_DEFAULT_CFG, forward_loop)
        # Keep the head's output layers, the bare convolutions and linears outside its Conv blocks, in float to limit
        # INT8 accuracy loss. DFL's fixed conv is left in float too.
        head = f"model.{len(model.model) - 1}."
        mtq.disable_quantizer(model, lambda n: n.startswith(head) and (".conv." not in n or ".dfl." in n))
    model.train(training)
    return model


def is_qat(model: nn.Module) -> bool:
    """Return True if the model carries fake-quantization modules inserted by `prepare_qat`.

    Matched by class name so that non-QAT models, i.e. every ordinary export, never import ModelOpt.
    """
    model = model.model if isinstance(getattr(model, "model", None), nn.Module) else model
    return any(type(m).__name__ == "TensorQuantizer" for m in model.modules())


def qat_state(model: nn.Module) -> dict[str, Any] | None:
    """Return the state that reproduces a model's fake-quantization, or None if it carries none.

    Ultralytics checkpoints are pickled modules, but ModelOpt builds its quantized layers as classes created at runtime,
    which pickle cannot look up on load. The quantization therefore travels beside the module as data. The checkpoint
    writers read it here and `restore_qat` reconstructs it at load and resume.

    Args:
        model (nn.Module): Model to read, left untouched.

    Returns:
        (dict | None): ModelOpt conversion state and the calibrated quantizer ranges, or None for a plain model.
    """
    model = getattr(model, "student_model", model)  # distillation checkpoints quantize only the student
    if not is_qat(model):
        return None
    import modelopt.torch.opt as mto

    return {
        "modelopt": mto.modelopt_state(model),
        "ranges": {k: v for k, v in model.state_dict().items() if "quantizer" in k},
    }


def strip_qat(model: nn.Module) -> None:
    """Revert a model's fake-quantization in place, leaving the plain layers it wraps.

    Checkpoint writers call this on the copy they are about to pickle, after `qat_state` has read the quantization out
    of it, since the runtime-generated layer classes cannot be pickled.
    """
    model = getattr(model, "student_model", model)  # distillation checkpoints quantize only the student
    if not is_qat(model):
        return
    from modelopt.torch.opt.conversion import ModeloptStateManager
    from modelopt.torch.opt.dynamic import DynamicModule

    for m in model.modules():
        if isinstance(m, DynamicModule):
            m.export()  # revert the runtime class to the plain layer it wraps
            m.__dict__.pop("_parallel_state", None)  # runtime process groups do not belong in a checkpoint
    ModeloptStateManager.remove_state(model)  # a reverted copy must not claim to be converted


def restore_qat(model: nn.Module, state: dict[str, Any]) -> None:
    """Re-apply the fake-quantization captured by `qat_state` to a model, in place."""
    model = getattr(model, "student_model", model)  # distillation checkpoints quantize only the student
    check_requirements(MODELOPT_REQUIREMENTS)
    import modelopt.torch.opt as mto

    mto.restore_from_modelopt_state(model, state["modelopt"])
    model.to(next(model.parameters()).device)
    model.load_state_dict(state["ranges"], strict=False)


def model_info(model, detailed=False, verbose=True, imgsz=640):
    """Print and return detailed model information layer by layer.

    Args:
        model (nn.Module): Model to analyze.
        detailed (bool, optional): Whether to print detailed layer information.
        verbose (bool, optional): Whether to print model information.
        imgsz (int | list | tuple, optional): Input image size, an int or an (h, w) pair.

    Returns:
        (tuple | None): Tuple containing the following, or None if `verbose` is False:
            - n_l (int): Number of layers.
            - n_p (int): Number of parameters.
            - n_g (int): Number of gradients.
            - flops (float): GFLOPs.
    """
    if not verbose:
        return
    n_p = get_num_params(model)  # number of parameters
    n_g = get_num_gradients(model)  # number of gradients
    layers = __import__("collections").OrderedDict((n, m) for n, m in model.named_modules() if len(m._modules) == 0)
    n_l = len(layers)  # number of layers
    if detailed:
        h = f"{'layer':>5}{'name':>40}{'type':>20}{'gradient':>10}{'parameters':>12}{'shape':>20}{'mu':>10}{'sigma':>10}"
        LOGGER.info(h)
        for i, (mn, m) in enumerate(layers.items()):
            mn = mn.replace("module_list.", "")
            mt = m.__class__.__name__
            if len(m._parameters):
                for pn, p in m.named_parameters():
                    LOGGER.info(
                        f"{i:>5g}{f'{mn}.{pn}':>40}{mt:>20}{p.requires_grad!r:>10}{p.numel():>12g}{list(p.shape)!s:>20}{p.mean():>10.3g}{p.std():>10.3g}{str(p.dtype).replace('torch.', ''):>15}"
                    )
            else:  # layers with no learnable params
                LOGGER.info(f"{i:>5g}{mn:>40}{mt:>20}{False!r:>10}{0:>12g}{[]!s:>20}{'-':>10}{'-':>10}{'-':>15}")

    flops = get_flops(model, imgsz)  # imgsz may be int or list, i.e. imgsz=640 or imgsz=[640, 320]
    fused = " (fused)" if getattr(model, "is_fused", lambda: False)() else ""
    fs = f", {flops:.1f} GFLOPs" if flops else ""
    yaml_file = getattr(model, "yaml_file", "") or getattr(model, "yaml", {}).get("yaml_file", "")
    model_name = Path(yaml_file).stem.replace("yolo", "YOLO") or "Model"
    LOGGER.info(f"{model_name} summary{fused}: {n_l:,} layers, {n_p:,} parameters, {n_g:,} gradients{fs}")
    return n_l, n_p, n_g, flops


def get_num_params(model):
    """Return the total number of parameters in a YOLO model."""
    return sum(x.numel() for x in model.parameters())


def get_num_gradients(model):
    """Return the total number of parameters with gradients in a YOLO model."""
    return sum(x.numel() for x in model.parameters() if x.requires_grad)


def model_info_for_loggers(trainer):
    """Return model info dict with useful model information.

    Args:
        trainer (ultralytics.engine.trainer.BaseTrainer): The trainer object containing model and validation data.

    Returns:
        (dict): Dictionary containing model parameters, GFLOPs, and inference speeds.

    Examples:
        YOLOv8n info for loggers
        >>> results = {
        ...     "model/parameters": 3151904,
        ...     "model/GFLOPs": 8.746,
        ...     "model/speed_ONNX(ms)": 41.244,
        ...     "model/speed_TensorRT(ms)": 3.211,
        ...     "model/speed_PyTorch(ms)": 18.755,
        ... }
    """
    if trainer.args.profile:  # profile ONNX and TensorRT times
        from ultralytics.utils.benchmarks import ProfileModels

        results = ProfileModels([trainer.last], device=trainer.device, imgsz=trainer.args.imgsz).run()[0]
        results.pop("model/name")
    else:  # only return PyTorch times from most recent validation
        results = {
            "model/parameters": get_num_params(trainer.model),
            "model/GFLOPs": round(get_flops(trainer.model, trainer.args.imgsz), 3),
        }
    results["model/speed_PyTorch(ms)"] = round(trainer.validator.speed["inference"], 3)
    return results


def get_flops(model, imgsz=640):
    """Calculate FLOPs (floating point operations) for a model in GFLOPs.

    Uses THOP's stride-aware image profiling, which extrapolates exactly from small stride-aligned proxy images, with
    the proxies widened past the anchor count at which an RT-DETR decoder's query selection saturates. Returns 0.0 if
    thop is unavailable or profiling fails.

    Args:
        model (nn.Module): The model to calculate FLOPs for.
        imgsz (int | list | tuple, optional): Input image size, an int or an (h, w) pair.

    Returns:
        (float): The model's GFLOPs (billions of floating point operations).
    """
    try:
        import thop
    except ImportError:
        thop = None  # conda support without 'ultralytics-thop' installed

    if not thop:
        return 0.0  # if not installed return 0.0 GFLOPs

    try:
        from ultralytics.nn.modules.head import RTDETRDecoder  # imported here: head.py imports this module

        model = unwrap_model(model)
        p = next(model.parameters())
        if not isinstance(imgsz, (list, tuple)):
            imgsz = [imgsz, imgsz]  # expand if int/float
        stride = max(int(model.stride.max()), 32) if hasattr(model, "stride") else 32
        # an RT-DETR decoder selects num_queries anchors, and a stride cell holds (4**nl - 1) / 3 anchors over nl levels
        min_cells = max(
            (-(-m.num_queries * 3 // (4**m.nl - 1)) for m in model.modules() if isinstance(m, RTDETRDecoder)), default=1
        )
        im = torch.empty((1, p.shape[1], *imgsz), device=p.device, dtype=p.dtype)  # input image in BCHW format
        return thop.profile(model, inputs=[im], stride=stride, min_cells=min_cells, verbose=False)[0] / 1e9 * 2
    except Exception:
        return 0.0


def initialize_weights(model):
    """Initialize model weights, biases, and module settings to default values."""
    for m in model.modules():
        t = type(m)
        if t is nn.Conv2d:
            pass  # nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
        elif t is nn.BatchNorm2d:
            m.eps = 1e-3
            m.momentum = 0.03


def scale_img(img, ratio=1.0, same_shape=False, gs=32):
    """Scale and pad an image tensor, optionally maintaining aspect ratio and padding to gs multiple.

    Args:
        img (torch.Tensor): Input image tensor with shape (B, C, H, W).
        ratio (float, optional): Scaling ratio.
        same_shape (bool, optional): Whether to pad or crop back to the original shape instead of a gs multiple.
        gs (int, optional): Grid size for padding.

    Returns:
        (torch.Tensor): Scaled and padded image tensor.
    """
    if ratio == 1.0:
        return img
    h, w = img.shape[2:]
    s = (int(h * ratio), int(w * ratio))  # new size
    img = F.interpolate(img, size=s, mode="bilinear", align_corners=False)  # resize
    if not same_shape:  # pad/crop img
        h, w = (math.ceil(x * ratio / gs) * gs for x in (h, w))
    return F.pad(img, [0, w - s[1], 0, h - s[0]], value=0.447)  # value = imagenet mean


def copy_attr(a, b, include=(), exclude=()):
    """Copy attributes from object 'b' to object 'a', with options to include/exclude certain attributes.

    Private attributes (names starting with '_') are never copied.

    Args:
        a (Any): Destination object to copy attributes to.
        b (Any): Source object to copy attributes from.
        include (tuple, optional): Attributes to include. If empty, all attributes are included.
        exclude (tuple, optional): Attributes to exclude.
    """
    for k, v in b.__dict__.items():
        if (len(include) and k not in include) or k.startswith("_") or k in exclude:
            continue
        else:
            setattr(a, k, v)


def intersect_dicts(da, db, exclude=()):
    """Return a dictionary of intersecting keys with matching shapes, excluding 'exclude' keys, using da values.

    Args:
        da (dict): First dictionary.
        db (dict): Second dictionary.
        exclude (tuple, optional): Keys to exclude.

    Returns:
        (dict): Dictionary of intersecting keys with matching shapes.
    """
    return {k: v for k, v in da.items() if k in db and all(x not in k for x in exclude) and v.shape == db[k].shape}


def is_parallel(model):
    """Return True if model is of type DP or DDP.

    Args:
        model (nn.Module): Model to check.

    Returns:
        (bool): True if model is DataParallel or DistributedDataParallel.
    """
    return isinstance(model, (nn.parallel.DataParallel, nn.parallel.DistributedDataParallel))


def unwrap_model(m: nn.Module) -> nn.Module:
    """Unwrap compiled and parallel models to get the base model.

    Args:
        m (nn.Module): A model that may be wrapped by torch.compile (._orig_mod) or parallel wrappers such as
            DataParallel/DistributedDataParallel (.module).

    Returns:
        (nn.Module): The unwrapped base model without compile or parallel wrappers.
    """
    while True:
        if hasattr(m, "_orig_mod") and isinstance(m._orig_mod, nn.Module):
            m = m._orig_mod
        elif hasattr(m, "module") and isinstance(m.module, nn.Module):
            m = m.module
        else:
            return m


def one_cycle(y1=0.0, y2=1.0, steps=100):
    """Return a lambda function for sinusoidal ramp from y1 to y2 https://arxiv.org/pdf/1812.01187.pdf.

    Args:
        y1 (float, optional): Initial value.
        y2 (float, optional): Final value.
        steps (int, optional): Number of steps.

    Returns:
        (function): Lambda function for computing the sinusoidal ramp.
    """
    return lambda x: max((1 - math.cos(x * math.pi / steps)) / 2, 0) * (y2 - y1) + y1


def init_seeds(seed=0, deterministic=False):
    """Initialize random number generator (RNG) seeds https://pytorch.org/docs/stable/notes/randomness.html.

    Args:
        seed (int, optional): Random seed.
        deterministic (bool, optional): Whether to set deterministic algorithms.
    """
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)  # also seeds every CUDA, MPS and XPU device
    # torch.backends.cudnn.benchmark = True  # AutoBatch problem https://github.com/ultralytics/yolov5/issues/9287
    if deterministic:
        if TORCH_2_0:
            torch.use_deterministic_algorithms(True, warn_only=True)  # warn if deterministic is not possible
            torch.backends.cudnn.deterministic = True
            if TORCH_2_2:  # skip deterministic mode's NaN fill of every new tensor, one fill kernel per allocation
                torch.utils.deterministic.fill_uninitialized_memory = False
            os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
            os.environ["PYTHONHASHSEED"] = str(seed)
        else:
            LOGGER.warning("Upgrade to torch>=2.0.0 for deterministic training.")
    else:
        unset_deterministic()


def unset_deterministic():
    """Unset all the configurations applied for deterministic training."""
    torch.use_deterministic_algorithms(False)
    torch.backends.cudnn.deterministic = False
    if TORCH_2_2:
        torch.utils.deterministic.fill_uninitialized_memory = True
    os.environ.pop("CUBLAS_WORKSPACE_CONFIG", None)
    os.environ.pop("PYTHONHASHSEED", None)


class ModelEMA:
    """Updated Exponential Moving Average (EMA) implementation.

    Keeps a moving average of everything in the model state_dict (parameters and buffers). For EMA details see
    References.

    To disable EMA set the `enabled` attribute to `False`.

    Attributes:
        ema (nn.Module): Copy of the model in evaluation mode.
        updates (int): Number of EMA updates.
        decay (function): Decay function that determines the EMA weight.
        enabled (bool): Whether EMA is enabled.

    References:
        - https://github.com/rwightman/pytorch-image-models
        - https://www.tensorflow.org/api_docs/python/tf/train/ExponentialMovingAverage
    """

    def __init__(self, model, decay=0.9999, tau=2000, updates=0):
        """Initialize EMA for 'model' with given arguments.

        Args:
            model (nn.Module): Model to create EMA for.
            decay (float, optional): Maximum EMA decay rate.
            tau (int, optional): EMA decay time constant.
            updates (int, optional): Initial number of updates.
        """
        self.ema = deepcopy(unwrap_model(model)).eval()  # FP32 EMA
        if hasattr(self.ema, "teacher_model"):
            # DistillationModel: strip the teacher so the EMA does not carry a full duplicate copy.
            self.ema.teacher_model = None
        self.updates = updates  # number of EMA updates
        self.decay = lambda x: decay * (1 - math.exp(-x / tau))  # decay exponential ramp (to help early epochs)
        for p in self.ema.parameters():
            p.requires_grad_(False)
        self.enabled = True
        self._pairs = None  # (ema tensors, model tensors) with floating dtype, built on the first update

    def update(self, model):
        """Update EMA parameters.

        Args:
            model (nn.Module): Model to update EMA from.
        """
        if self.enabled:
            self.updates += 1
            d = self.decay(self.updates)

            if self._pairs is None:  # the tensors are updated in place, so the lists are built once
                msd = unwrap_model(model).state_dict()  # model state_dict
                ema_v, model_v = [], []
                for k, v in self.ema.state_dict().items():
                    if v.dtype.is_floating_point:  # true for FP16 and FP32
                        ema_v.append(v)
                        model_v.append(msd[k])
                self._pairs = ema_v, model_v
            ema_v, model_v = self._pairs
            if (
                ema_v and TORCH_2_0 and ema_v[0].device.type != "npu" and (TORCH_2_4 or ema_v[0].device.type != "mps")
            ):  # one kernel launch per op
                torch._foreach_lerp_(ema_v, model_v, 1 - d)
            else:  # _foreach_lerp_ needs torch>=2.0, MPS torch>=2.4, and is unavailable on NPU
                for v, m in zip(ema_v, model_v):
                    v.mul_(d).add_(m, alpha=1 - d)

    def update_attr(self, model, include=(), exclude=()):
        """Copy attributes from model to EMA, with options to include/exclude certain attributes.

        Args:
            model (nn.Module): Model to copy attributes from; compile and parallel wrappers are unwrapped first.
            include (tuple, optional): Attributes to include.
            exclude (tuple, optional): Attributes to exclude.
        """
        if self.enabled:
            copy_attr(self.ema, unwrap_model(model), include, exclude)


def strip_optimizer(f: str | Path = "best.pt", s: str = "", updates: dict[str, Any] | None = None) -> dict[str, Any]:
    """Strip optimizer from 'f' to finalize training, optionally save as 's'.

    Args:
        f (str | Path): File path to model to strip the optimizer from.
        s (str, optional): File path to save the model with stripped optimizer to. If not provided, 'f' will be
            overwritten.
        updates (dict, optional): A dictionary of updates to overlay onto the checkpoint before saving.

    Returns:
        (dict): The combined checkpoint dictionary, or an empty dict if 'f' is not a valid Ultralytics checkpoint.

    Examples:
        >>> from pathlib import Path
        >>> from ultralytics.utils.torch_utils import strip_optimizer
        >>> for f in Path("path/to/model/checkpoints").rglob("*.pt"):
        ...     strip_optimizer(f)
    """
    try:
        x = torch_load(f, map_location=torch.device("cpu"))
        assert isinstance(x, dict), "checkpoint is not a Python dictionary"
        assert "model" in x, "'model' missing from checkpoint"
    except Exception as e:
        LOGGER.warning(f"Skipping {f}, not a valid Ultralytics model: {e}")
        return {}

    metadata = {
        "date": datetime.now().astimezone().isoformat(),
        "version": __version__,
        "license": "AGPL-3.0 License (https://ultralytics.com/license)",
        "docs": "https://docs.ultralytics.com",
    }

    # Update model
    if x.get("ema"):
        x["model"] = x["ema"]  # replace model with EMA

    # Unwrap DistillationModel to save only the student model
    from ultralytics.nn.distill_model import DistillationModel

    if isinstance(x["model"], DistillationModel):
        x["model"]._remove_feature_hooks()
        x["model"] = x["model"].student_model

    if hasattr(x["model"], "args"):
        x["model"].args = dict(x["model"].args)  # convert from IterableSimpleNamespace to dict
    if hasattr(x["model"], "criterion"):
        x["model"].criterion = None  # strip loss criterion
    x["model"].half()  # to FP16
    for p in x["model"].parameters():
        p.requires_grad = False

    # Update other keys
    args = {**DEFAULT_CFG_DICT, **x.get("train_args", {})}  # combine args
    for k in "optimizer", "best_fitness", "stopper", "ema", "updates", "scaler":  # keys
        x[k] = None
    x["epoch"] = -1
    x["train_args"] = {k: v for k, v in args.items() if k in DEFAULT_CFG_KEYS}  # strip non-default keys
    # x['model'].args = x['train_args']

    # Save
    combined = {**metadata, **x, **(updates or {})}
    torch.save(combined, s or f)  # combine dicts (prefer to the right)
    mb = os.path.getsize(s or f) / 1e6  # file size
    LOGGER.info(f"Optimizer stripped from {f},{f' saved as {s},' if s else ''} {mb:.1f}MB")
    return combined


def convert_optimizer_state_dict_to_fp16(state_dict):
    """Convert the state_dict of a given optimizer to FP16, focusing on the 'state' key for tensor conversions.

    FP32 tensors are converted in place, except 'step' and 'exp_avg_sq', which are kept in FP32.

    Args:
        state_dict (dict): Optimizer state dictionary.

    Returns:
        (dict): Converted optimizer state dictionary with FP16 tensors.
    """
    for state in state_dict["state"].values():
        for k, v in state.items():
            if k not in {"step", "exp_avg_sq"} and isinstance(v, torch.Tensor) and v.dtype is torch.float32:
                state[k] = v.half()

    return state_dict


@contextmanager
def cuda_memory_usage(device=None):
    """Monitor and manage accelerator memory usage.

    This function empties the active accelerator cache, yields a dictionary containing memory usage information, and
    then records the reserved memory on the specified device.

    Args:
        device (torch.device, optional): The accelerator device to query memory usage for. CPU and MPS devices are not
            measured.

    Yields:
        (dict): A dictionary with a key 'memory' initialized to 0, updated with reserved memory in bytes.
    """
    info = {"memory": 0}
    if device is not None and device.type in {"cpu", "mps"}:
        yield info
        return
    accelerator = get_torch_device_backend(device or "cuda")
    if accelerator.is_available() and hasattr(accelerator, "memory_reserved"):
        accelerator.empty_cache()
        try:
            yield info
        finally:
            info["memory"] = accelerator.memory_reserved(device)
    else:
        yield info


def profile_ops(input, ops, n=10, device=None, max_num_obj=0):
    """Ultralytics speed, memory and FLOPs profiler.

    Args:
        input (torch.Tensor | list): Input tensor(s) to profile.
        ops (nn.Module | Callable | list): Model, callable, or list of operations to profile.
        n (int, optional): Number of iterations to average.
        device (str | torch.device, optional): Device to profile on.
        max_num_obj (int, optional): Maximum number of objects per image used to simulate training-loss memory for
            AutoBatch. Requires `ops` to have `stride` and `names` attributes. 0 disables the simulation.

    Returns:
        (list): Profile results for each input and operation pair, each either a list of [parameters, GFLOPs, memory
            (GB), forward time (ms), backward time (ms), input shape, output shape] or None if profiling failed.

    Examples:
        >>> from ultralytics.utils.torch_utils import profile_ops
        >>> input = torch.randn(16, 3, 640, 640)
        >>> m1 = lambda x: x * torch.sigmoid(x)
        >>> m2 = nn.SiLU()
        >>> profile_ops(input, [m1, m2], n=100)  # profile over 100 iterations
    """
    try:
        import thop
    except ImportError:
        thop = None  # conda support without 'ultralytics-thop' installed

    results = []
    device = select_device(device, verbose=False)
    LOGGER.info(
        f"{'Params':>12s}{'GFLOPs':>12s}{'GPU_mem (GB)':>14s}{'forward (ms)':>14s}{'backward (ms)':>14s}"
        f"{'input':>24s}{'output':>24s}"
    )
    gc.collect()  # attempt to free unused memory
    accelerator = get_torch_device_backend(device) if device.type not in {"cpu", "mps"} else None
    if accelerator is not None:
        accelerator.empty_cache()
    for x in input if isinstance(input, list) else [input]:
        x = x.to(device)
        x.requires_grad = True
        for m in ops if isinstance(ops, list) else [ops]:
            m = m.to(device) if hasattr(m, "to") else m  # device
            m = m.half() if hasattr(m, "half") and isinstance(x, torch.Tensor) and x.dtype is torch.float16 else m
            tf, tb, t = 0, 0, [0, 0, 0]  # dt forward, backward
            try:
                flops = thop.profile(m, inputs=[x], verbose=False)[0] / 1e9 * 2 if thop else 0  # GFLOPs
            except Exception:
                flops = 0

            try:
                mem = 0
                for _ in range(n):
                    with cuda_memory_usage(device) as cuda_info:
                        t[0] = time_sync(device)
                        y = m(x)
                        t[1] = time_sync(device)
                        try:
                            sum(yi.sum() for yi in tree_flatten(y)[0] if isinstance(yi, torch.Tensor)).backward()
                            t[2] = time_sync(device)
                        except Exception:  # no backward method
                            # print(e)  # for debug
                            t[2] = float("nan")
                    mem += cuda_info["memory"] / 1e9  # (GB)
                    tf += (t[1] - t[0]) * 1000 / n  # ms per op forward
                    tb += (t[2] - t[1]) * 1000 / n  # ms per op backward
                    if max_num_obj:  # simulate training with predictions per image grid (for AutoBatch)
                        with cuda_memory_usage(device) as cuda_info:
                            anchors = int(sum((x.shape[-1] / s) * (x.shape[-2] / s) for s in m.stride.tolist()))
                            # Conservative detect-loss envelope: ~6 fp32-equivalents each for TaskAlignedAssigner
                            # metric/top-k state and the cls path (pred/target + two op temps of unreduced BCE:
                            # ~4 in pure fp32, ~6 under AMP where autocast upcasts both BCE inputs to fp32 copies)
                            sim = (
                                torch.randn(x.shape[0], 6 * max_num_obj, anchors, device=device, dtype=torch.float32),
                                torch.randn(x.shape[0], anchors, 6 * len(m.names), device=device, dtype=torch.float32),
                            )
                        del sim
                        mem += cuda_info["memory"] / 1e9  # (GB)
                s_in, s_out = (tuple(x.shape) if isinstance(x, torch.Tensor) else "list" for x in (x, y))  # shapes
                p = sum(x.numel() for x in m.parameters()) if isinstance(m, nn.Module) else 0  # parameters
                LOGGER.info(f"{p:12}{flops:12.4g}{mem:>14.3f}{tf:14.4g}{tb:14.4g}{s_in!s:>24s}{s_out!s:>24s}")
                results.append([p, flops, mem, tf, tb, s_in, s_out])
            except Exception as e:
                LOGGER.info(e)
                results.append(None)
            finally:
                gc.collect()  # attempt to free unused memory
                if accelerator is not None:
                    accelerator.empty_cache()
    return results


class EarlyStopping:
    """Early stopping class that stops training when a specified number of epochs have passed without improvement.

    Attributes:
        best_fitness (float): Best fitness value observed.
        best_epoch (int): Epoch where best fitness was observed.
        patience (int): Number of epochs to wait after fitness stops improving before stopping.
        possible_stop (bool): Flag indicating if stopping may occur next epoch.
    """

    def __init__(self, patience=50):
        """Initialize early stopping object.

        Args:
            patience (int, optional): Number of epochs to wait after fitness stops improving before stopping. 0 or None
                disables early stopping. The trainer always passes the cfg `patience` value (100 by default).
        """
        self.best_fitness = 0.0  # i.e. mAP
        self.best_epoch = 0
        self.patience = patience or float("inf")  # epochs to wait after fitness stops improving to stop
        self.possible_stop = False  # possible stop may occur next epoch

    def __call__(self, epoch, fitness):
        """Check whether to stop training.

        Args:
            epoch (int): Current epoch of training.
            fitness (float): Fitness value of current epoch.

        Returns:
            (bool): True if training should stop, False otherwise.
        """
        if fitness is None:  # check if fitness=None (happens when val=False)
            return False

        if fitness > self.best_fitness or self.best_fitness == 0:  # allow for early zero-fitness stage of training
            self.best_epoch = epoch
            self.best_fitness = fitness
        delta = epoch - self.best_epoch  # epochs without improvement
        self.possible_stop = delta >= (self.patience - 1)  # possible stop may occur next epoch
        stop = delta >= self.patience  # stop training if patience exceeded
        if stop:
            prefix = colorstr("EarlyStopping: ")
            LOGGER.info(
                f"{prefix}Training stopped early as no improvement observed in last {self.patience} epochs. "
                f"Best results observed at epoch {self.best_epoch}, best model saved as best.pt.\n"
                f"To update EarlyStopping(patience={self.patience}) pass a new patience value, "
                f"i.e. `patience=300` or use `patience=0` to disable EarlyStopping."
            )
        return stop


def attempt_compile(
    model: torch.nn.Module,
    device: torch.device,
    imgsz: int = 640,
    use_autocast: bool = False,
    warmup: bool = False,
    mode: bool | str = "default",
) -> torch.nn.Module:
    """Compile a model with torch.compile and optionally warm up the graph to reduce first-iteration latency.

    This utility attempts to compile the provided model using the inductor backend. If compilation is unavailable or
    fails, the original model is returned unchanged. An optional warmup performs a single forward pass on a dummy input
    to prime the compiled graph and measure compile/warmup time.

    Args:
        model (torch.nn.Module): Model to compile.
        device (torch.device): Inference device used for warmup and autocast decisions.
        imgsz (int, optional): Square input size to create a dummy tensor with shape (1, 3, imgsz, imgsz) for warmup.
        use_autocast (bool, optional): Whether to run warmup under autocast on CUDA or MPS devices.
        warmup (bool, optional): Whether to execute a single dummy forward pass to warm up the compiled model.
        mode (bool | str, optional): torch.compile mode. True → "default", False → no compile, or a string like
            "default", "reduce-overhead", "max-autotune-no-cudagraphs".

    Returns:
        (torch.nn.Module): Compiled model if compilation succeeds, otherwise the original unmodified model.

    Raises:
        ValueError: If the model carries QAT fake-quantization modules, which do not support torch.compile.

    Examples:
        >>> import torch
        >>> from ultralytics.utils.torch_utils import attempt_compile
        >>> device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
        >>> model = torch.nn.Conv2d(3, 16, 3).to(device)
        >>> # Try to compile and warm up a model with a 640x640 input
        >>> model = attempt_compile(model, device=device, imgsz=640, use_autocast=True, warmup=True)

    Notes:
        - If the current PyTorch build does not provide torch.compile or `mode` is False, the function returns the input
          model immediately.
        - Compilation is lazy and runs at the first forward pass, so the inductor CPU prerequisite of a host C++
          compiler is verified up front and the original model is returned if none is available.
        - Warmup runs under torch.inference_mode and may use torch.autocast for CUDA/MPS to align compute precision.
        - CUDA devices are synchronized after warmup to account for asynchronous kernel execution.
    """
    if not hasattr(torch, "compile") or not mode:
        return model
    if is_qat(model):
        raise ValueError("QAT models do not support torch.compile. Use compile=False.")

    if mode is True:
        mode = "default"
    prefix = colorstr("compile:")
    if device.type == "cpu":
        try:  # compilation is lazy, so verify the inductor CPU requirement of a host C++ compiler before compiling
            from torch._inductor.cpp_builder import get_cpp_compiler

            get_cpp_compiler()
        except ImportError:
            pass  # older torch without cpp_builder, defer to torch.compile
        except Exception as e:
            LOGGER.warning(f"{prefix} no C++ compiler found for the inductor backend, continuing uncompiled: {e}")
            return model
    LOGGER.info(f"{prefix} starting torch.compile with '{mode}' mode...")
    t0 = time.perf_counter()
    try:
        model = torch.compile(model, mode=mode, backend="inductor")
    except Exception as e:
        LOGGER.warning(f"{prefix} torch.compile failed, continuing uncompiled: {e}")
        return model
    t_compile = time.perf_counter() - t0

    t_warm = 0.0
    if warmup:
        # Use a single dummy tensor to build the graph shape state and reduce first-iteration latency
        dummy = torch.zeros(1, 3, imgsz, imgsz, device=device)
        if use_autocast and device.type == "cuda":
            dummy = dummy.half()
        t1 = time.perf_counter()
        with torch.inference_mode():
            if use_autocast and device.type in {"cuda", "mps"}:
                with torch.autocast(device.type):
                    _ = model(dummy)
            else:
                _ = model(dummy)
        if device.type == "cuda":
            torch.cuda.synchronize(device)
        t_warm = time.perf_counter() - t1

    total = t_compile + t_warm
    if warmup:
        LOGGER.info(f"{prefix} complete in {total:.1f}s (compile {t_compile:.1f}s + warmup {t_warm:.1f}s)")
    else:
        LOGGER.info(f"{prefix} compile complete in {t_compile:.1f}s (no warmup)")
    return model
