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

from __future__ import annotations

import contextlib
import math
import re
import time

import cv2
import numpy as np
import torch
import torch.nn.functional as F

from ultralytics.utils import NOT_MACOS14, TORCH_VERSION
from ultralytics.utils.checks import check_version
from ultralytics.utils.torch_utils import get_torch_device_backend


class Profile(contextlib.ContextDecorator):
    """Ultralytics Profile class for timing code execution.

    Use as a decorator with @Profile() or as a context manager with 'with Profile():'. Provides accurate timing
    measurements with accelerator synchronization support.

    Attributes:
        t (float): Accumulated time in seconds.
        dt (float): Elapsed time in seconds of the most recent timed block.
        device (torch.device): Device used for model inference.
        accelerator (module | None): PyTorch device module used for timing synchronization, None if not needed.

    Examples:
        Use as a context manager to time code execution
        >>> with Profile() as dt:
        ...     pass  # slow operation here
        >>> str(dt).startswith("Elapsed time is ")
        True

        Use as a decorator to time function execution
        >>> @Profile()
        ... def slow_function():
        ...     time.sleep(0.1)
    """

    def __init__(self, t: float = 0.0, device: torch.device | None = None):
        """Initialize the Profile class.

        Args:
            t (float): Initial accumulated time in seconds.
            device (torch.device, optional): Device used for model inference to enable accelerator synchronization.
        """
        self.t = t
        self.device = device
        device_type = getattr(device, "type", str(device).split(":")[0] if device else None)
        self.accelerator = get_torch_device_backend(device_type) if device_type in {"cuda", "npu", "xpu"} else None

    def __enter__(self):
        """Start timing."""
        self.start = self.time()
        return self

    def __exit__(self, type, value, traceback):
        """Stop timing."""
        self.dt = self.time() - self.start  # delta-time
        self.t += self.dt  # accumulate dt

    def __str__(self):
        """Return a human-readable string representing the accumulated elapsed time."""
        return f"Elapsed time is {self.t} s"

    def time(self):
        """Return the current time with accelerator synchronization if applicable."""
        if self.accelerator is not None:
            self.accelerator.synchronize(self.device)
        return time.perf_counter()


def segment2box(segment: np.ndarray, width: int = 640, height: int = 640) -> np.ndarray:
    """Convert segment coordinates to bounding box coordinates.

    Converts a single segment label to a box label by finding the minimum and maximum x and y coordinates of the polygon
    clipped to the image, so segments crossing the image boundary keep their visible extent. Segments entirely inside
    the image, or whose bounding box lies entirely outside it, return immediately without clipping.

    Args:
        segment (np.ndarray): Segment coordinates in format (N, 2) where N is number of points.
        width (int): Width of the image in pixels.
        height (int): Height of the image in pixels.

    Returns:
        (np.ndarray): Bounding box coordinates in xyxy format [x1, y1, x2, y2], or zeros if the segment is empty or
            entirely outside the image.
    """
    if not len(segment):
        return np.zeros(4, dtype=segment.dtype)
    x, y = segment[:, 0], segment[:, 1]
    xmin, ymin, xmax, ymax = x.min(), y.min(), x.max(), y.max()
    if xmin >= 0 and ymin >= 0 and xmax <= width and ymax <= height:  # fully inside image
        return np.array([xmin, ymin, xmax, ymax], dtype=segment.dtype)
    if xmax < 0 or ymax < 0 or xmin > width or ymin > height:  # fully outside image
        return np.zeros(4, dtype=segment.dtype)
    axes = np.array((0, 0, 1, 1))
    bounds = np.array((0, width, 0, height), dtype=segment.dtype)
    lims = np.array((height, height, width, width), dtype=segment.dtype)  # (height, width)[axis] per boundary
    start, delta = segment, np.roll(segment, -1, axis=0) - segment
    with np.errstate(divide="ignore", invalid="ignore"):
        t = (bounds - start[:, axes]) / delta[:, axes]
        inter = start[:, None, :] + t[:, :, None] * delta[:, None, :]
    other = inter[:, np.arange(4), 1 - axes]
    corners = np.array(((0, 0), (width, 0), (0, height), (width, height)), dtype=segment.dtype)
    contour = segment.astype(np.float32)
    points = np.concatenate(
        (
            segment[(x >= 0) & (y >= 0) & (x <= width) & (y <= height)],
            inter[(t >= 0) & (t <= 1) & (other >= 0) & (other <= lims)],
            corners[[cv2.pointPolygonTest(contour, tuple(map(float, p)), False) >= 0 for p in corners]],
        )
    )
    return (
        np.array([*points.min(0), *points.max(0)], dtype=segment.dtype)
        if len(points)
        else np.zeros(4, dtype=segment.dtype)
    )


def scale_boxes(
    img1_shape: tuple[int, int],
    boxes: torch.Tensor | np.ndarray,
    img0_shape: tuple[int, int],
    ratio_pad: tuple | None = None,
    padding: bool = True,
    xywh: bool = False,
) -> torch.Tensor | np.ndarray:
    """Rescale bounding boxes from one image shape to another.

    Rescales bounding boxes from img1_shape to img0_shape, accounting for padding and aspect ratio changes. Supports
    both xyxy and xywh box formats. Boxes are modified in place.

    Args:
        img1_shape (tuple[int, int]): Shape of the source image (height, width).
        boxes (torch.Tensor | np.ndarray): Bounding boxes to rescale in format (N, 4).
        img0_shape (tuple[int, int]): Shape of the target image (height, width).
        ratio_pad (tuple, optional): Ratio and padding as ((ratio_h, ratio_w), (pad_w, pad_h)).
        padding (bool): Whether boxes are based on YOLO-style augmented images with padding.
        xywh (bool): Whether box format is xywh (True) or xyxy (False).

    Returns:
        (torch.Tensor | np.ndarray): Rescaled bounding boxes in the same format as input, clipped to img0_shape for xyxy
            boxes.
    """
    if ratio_pad is None:  # calculate from img0_shape
        gain = min(img1_shape[0] / img0_shape[0], img1_shape[1] / img0_shape[1])  # gain  = old / new
        new_h, new_w = round(img0_shape[0] * gain), round(img0_shape[1] * gain)  # LetterBox rounds each side
        gain_y, gain_x = new_h / img0_shape[0], new_w / img0_shape[1]
        pad_x, pad_y = round((img1_shape[1] - new_w) / 2 - 0.1), round((img1_shape[0] - new_h) / 2 - 0.1)
    else:
        gain_y, gain_x = ratio_pad[0]
        pad_x, pad_y = ratio_pad[1]

    if padding:
        boxes[..., 0] -= pad_x  # x padding
        boxes[..., 1] -= pad_y  # y padding
        if not xywh:
            boxes[..., 2] -= pad_x  # x padding
            boxes[..., 3] -= pad_y  # y padding
    boxes[..., 0] /= gain_x
    boxes[..., 1] /= gain_y
    boxes[..., 2] /= gain_x
    boxes[..., 3] /= gain_y
    return boxes if xywh else clip_boxes(boxes, img0_shape)


def make_divisible(x: float, divisor):
    """Return the smallest number >= x that is divisible by the given divisor.

    Args:
        x (int | float): The number to make divisible.
        divisor (int | torch.Tensor): The divisor.

    Returns:
        (int): The smallest number >= x divisible by the divisor.
    """
    if isinstance(divisor, torch.Tensor):
        divisor = int(divisor.max())  # to int
    return math.ceil(x / divisor) * divisor


def clip_boxes(boxes, shape):
    """Clip bounding boxes to image boundaries in place.

    Args:
        boxes (torch.Tensor | np.ndarray): Bounding boxes in xyxy format to clip.
        shape (tuple): Image shape as HWC or HW (supports both).

    Returns:
        (torch.Tensor | np.ndarray): Clipped bounding boxes.
    """
    h, w = shape[:2]  # supports both HWC or HW shapes
    if isinstance(boxes, torch.Tensor):  # faster individually
        if NOT_MACOS14 and not (boxes.device.type == "mps" and check_version(TORCH_VERSION, "<2.5.0")):
            boxes[..., 0].clamp_(0, w)  # x1
            boxes[..., 1].clamp_(0, h)  # y1
            boxes[..., 2].clamp_(0, w)  # x2
            boxes[..., 3].clamp_(0, h)  # y2
        else:  # MPS strided in-place bug on macOS 14 or torch<2.5
            boxes[..., 0] = boxes[..., 0].clamp(0, w)
            boxes[..., 1] = boxes[..., 1].clamp(0, h)
            boxes[..., 2] = boxes[..., 2].clamp(0, w)
            boxes[..., 3] = boxes[..., 3].clamp(0, h)
    else:  # np.array (faster grouped)
        boxes[..., [0, 2]] = boxes[..., [0, 2]].clip(0, w)  # x1, x2
        boxes[..., [1, 3]] = boxes[..., [1, 3]].clip(0, h)  # y1, y2
    return boxes


def clip_coords(coords, shape):
    """Clip line coordinates to image boundaries in place.

    Args:
        coords (torch.Tensor | np.ndarray): Line coordinates to clip, with x and y in the first two channels of the last
            dimension.
        shape (tuple): Image shape as HWC or HW (supports both).

    Returns:
        (torch.Tensor | np.ndarray): Clipped coordinates.
    """
    h, w = shape[:2]  # supports both HWC or HW shapes
    if isinstance(coords, torch.Tensor):
        if NOT_MACOS14 and not (coords.device.type == "mps" and check_version(TORCH_VERSION, "<2.5.0")):
            coords[..., 0].clamp_(0, w)  # x
            coords[..., 1].clamp_(0, h)  # y
        else:  # MPS strided in-place bug on macOS 14 or torch<2.5
            coords[..., 0] = coords[..., 0].clamp(0, w)
            coords[..., 1] = coords[..., 1].clamp(0, h)
    else:  # np.array
        coords[..., 0] = coords[..., 0].clip(0, w)  # x
        coords[..., 1] = coords[..., 1].clip(0, h)  # y
    return coords


def xyxy2xywh(x):
    """Convert bounding box coordinates from (x1, y1, x2, y2) format to (x, y, width, height) format.

    (x1, y1) is the top-left corner, (x2, y2) is the bottom-right corner, and (x, y) is the box center.

    Args:
        x (np.ndarray | torch.Tensor | list | tuple): Input bounding box coordinates in (x1, y1, x2, y2) format.

    Returns:
        (np.ndarray | torch.Tensor): Bounding box coordinates in (x, y, width, height) format.
    """
    if isinstance(x, (list, tuple)):
        x = np.asarray(x, dtype=np.float32)  # float so odd integer boxes keep fractional centers
    assert x.shape[-1] == 4, f"input shape last dimension expected 4 but input shape is {x.shape}"
    y = empty_like(x)  # faster than clone/copy
    x1, y1, x2, y2 = x[..., 0], x[..., 1], x[..., 2], x[..., 3]
    y[..., 0] = (x1 + x2) / 2  # x center
    y[..., 1] = (y1 + y2) / 2  # y center
    y[..., 2] = x2 - x1  # width
    y[..., 3] = y2 - y1  # height
    return y


def xywh2xyxy(x):
    """Convert bounding box coordinates from (x, y, width, height) format to (x1, y1, x2, y2) format.

    (x, y) is the box center, (x1, y1) is the top-left corner, and (x2, y2) is the bottom-right corner. Note: ops per 2
    channels faster than per channel.

    Args:
        x (np.ndarray | torch.Tensor | list | tuple): Input bounding box coordinates in (x, y, width, height) format.

    Returns:
        (np.ndarray | torch.Tensor): Bounding box coordinates in (x1, y1, x2, y2) format.
    """
    if isinstance(x, (list, tuple)):
        x = np.asarray(x, dtype=np.float32)  # float so odd integer boxes keep fractional centers
    assert x.shape[-1] == 4, f"input shape last dimension expected 4 but input shape is {x.shape}"
    y = empty_like(x)  # faster than clone/copy
    xy = x[..., :2]  # centers
    wh = x[..., 2:] / 2  # half width-height
    y[..., :2] = xy - wh  # top left xy
    y[..., 2:] = xy + wh  # bottom right xy
    return y


def xywhn2xyxy(x, w: int = 640, h: int = 640, padw: int = 0, padh: int = 0):
    """Convert normalized bounding box coordinates to pixel coordinates.

    Args:
        x (np.ndarray | torch.Tensor): Normalized bounding box coordinates in (x, y, w, h) format.
        w (int): Image width in pixels.
        h (int): Image height in pixels.
        padw (int): Horizontal padding in pixels added to x coordinates.
        padh (int): Vertical padding in pixels added to y coordinates.

    Returns:
        (np.ndarray | torch.Tensor): Bounding box coordinates in (x1, y1, x2, y2) format.
    """
    assert x.shape[-1] == 4, f"input shape last dimension expected 4 but input shape is {x.shape}"
    y = empty_like(x)  # faster than clone/copy
    xc, yc, xw, xh = x[..., 0], x[..., 1], x[..., 2], x[..., 3]
    half_w, half_h = xw / 2, xh / 2
    y[..., 0] = w * (xc - half_w) + padw  # top left x
    y[..., 1] = h * (yc - half_h) + padh  # top left y
    y[..., 2] = w * (xc + half_w) + padw  # bottom right x
    y[..., 3] = h * (yc + half_h) + padh  # bottom right y
    return y


def xyxy2xywhn(x, w: int = 640, h: int = 640, clip: bool = False, eps: float = 0.0):
    """Convert bounding box coordinates from (x1, y1, x2, y2) format to normalized (x, y, width, height) format.

    x, y, width and height are normalized to image dimensions.

    Args:
        x (np.ndarray | torch.Tensor): Input bounding box coordinates in (x1, y1, x2, y2) format.
        w (int): Image width in pixels.
        h (int): Image height in pixels.
        clip (bool): Whether to clip boxes to image boundaries (in place) before conversion.
        eps (float): Margin subtracted from image width and height when clipping.

    Returns:
        (np.ndarray | torch.Tensor): Normalized bounding box coordinates in (x, y, width, height) format.
    """
    if clip:
        x = clip_boxes(x, (h - eps, w - eps))
    assert x.shape[-1] == 4, f"input shape last dimension expected 4 but input shape is {x.shape}"
    y = empty_like(x)  # faster than clone/copy
    x1, y1, x2, y2 = x[..., 0], x[..., 1], x[..., 2], x[..., 3]
    y[..., 0] = ((x1 + x2) / 2) / w  # x center
    y[..., 1] = ((y1 + y2) / 2) / h  # y center
    y[..., 2] = (x2 - x1) / w  # width
    y[..., 3] = (y2 - y1) / h  # height
    return y


def xywh2ltwh(x):
    """Convert bounding box format from [x, y, w, h] to [x1, y1, w, h] where x1, y1 are top-left coordinates.

    Args:
        x (np.ndarray | torch.Tensor): Input bounding box coordinates in xywh format.

    Returns:
        (np.ndarray | torch.Tensor): Bounding box coordinates in ltwh format.
    """
    y = x.clone() if isinstance(x, torch.Tensor) else np.copy(x)
    y[..., 0] = x[..., 0] - x[..., 2] / 2  # top left x
    y[..., 1] = x[..., 1] - x[..., 3] / 2  # top left y
    return y


def xyxy2ltwh(x):
    """Convert bounding boxes from [x1, y1, x2, y2] to [x1, y1, w, h] format.

    Args:
        x (np.ndarray | torch.Tensor): Input bounding box coordinates in xyxy format.

    Returns:
        (np.ndarray | torch.Tensor): Bounding box coordinates in ltwh format.
    """
    y = x.clone() if isinstance(x, torch.Tensor) else np.copy(x)
    y[..., 2] = x[..., 2] - x[..., 0]  # width
    y[..., 3] = x[..., 3] - x[..., 1]  # height
    return y


def ltwh2xywh(x):
    """Convert bounding boxes from [x1, y1, w, h] to [x, y, w, h] where xy1=top-left, xy=center.

    Args:
        x (np.ndarray | torch.Tensor): Input bounding box coordinates.

    Returns:
        (np.ndarray | torch.Tensor): Bounding box coordinates in xywh format.
    """
    y = x.clone() if isinstance(x, torch.Tensor) else np.copy(x)
    y[..., 0] = x[..., 0] + x[..., 2] / 2  # center x
    y[..., 1] = x[..., 1] + x[..., 3] / 2  # center y
    return y


def xyxyxyxy2xywhr(x):
    """Convert batched Oriented Bounding Boxes (OBB) from [xy1, xy2, xy3, xy4] to [xywh, rotation] format.

    Args:
        x (np.ndarray | torch.Tensor): Input box corners with shape (N, 8) or (N, 4, 2) in [xy1, xy2, xy3, xy4] format.
            Polygons with more than four points are accepted in the same two layouts, (N, 2P) or (N, P, 2), and are
            reduced to their minimum-area rectangle.

    Returns:
        (np.ndarray | torch.Tensor): Converted data in [cx, cy, w, h, rotation] format with shape (N, 5). The
            parameterization is canonical rather than the caller's: w is the longer side and rotation is in radians
            from [-pi/4, 3pi/4), so a box given with w < h comes back with w and h swapped and its angle shifted by
            pi/2 modulo pi.
    """
    is_torch = isinstance(x, torch.Tensor)
    points = x.cpu().numpy() if is_torch else x
    rboxes = []
    for pts in points:
        # NOTE: Use cv2.minAreaRect to get accurate xywhr,
        # especially some objects are cut off by augmentations in dataloader.
        (cx, cy), (w, h), angle = cv2.minAreaRect(pts.reshape(-1, 2))
        # convert angle to radian and normalize to [-pi/4, 3pi/4)
        theta = angle / 180 * np.pi
        if w < h:
            w, h = h, w
            theta += np.pi / 2
        while theta >= 3 * np.pi / 4:
            theta -= np.pi
        while theta < -np.pi / 4:
            theta += np.pi
        rboxes.append([cx, cy, w, h, theta])
    rboxes = np.asarray(rboxes).reshape(-1, 5)  # reshape keeps the (0, 5) shape on an empty input
    return torch.tensor(rboxes, device=x.device, dtype=x.dtype) if is_torch else rboxes


def xywhr2xyxyxyxy(x):
    """Convert batched Oriented Bounding Boxes (OBB) from [xywh, rotation] to [xy1, xy2, xy3, xy4] format.

    Args:
        x (np.ndarray | torch.Tensor): Boxes in [cx, cy, w, h, rotation] format with shape (N, 5) or (B, N, 5). Rotation
            is in radians and is neither range-checked nor normalized; the box is not canonicalized, so converting the
            (N, 4, 2) corners back with xyxyxyxy2xywhr returns the canonical form of the same rectangle rather than
            these values.

    Returns:
        (np.ndarray | torch.Tensor): Converted corner points with shape (N, 4, 2) or (B, N, 4, 2).
    """
    cos, sin, cat, stack = (
        (torch.cos, torch.sin, torch.cat, torch.stack)
        if isinstance(x, torch.Tensor)
        else (np.cos, np.sin, np.concatenate, np.stack)
    )

    ctr = x[..., :2]
    w, h, angle = (x[..., i : i + 1] for i in range(2, 5))
    cos_value, sin_value = cos(angle), sin(angle)
    vec1 = [w / 2 * cos_value, w / 2 * sin_value]
    vec2 = [-h / 2 * sin_value, h / 2 * cos_value]
    vec1 = cat(vec1, -1)
    vec2 = cat(vec2, -1)
    pt1 = ctr + vec1 + vec2
    pt2 = ctr + vec1 - vec2
    pt3 = ctr - vec1 - vec2
    pt4 = ctr - vec1 + vec2
    return stack([pt1, pt2, pt3, pt4], -2)


def ltwh2xyxy(x):
    """Convert bounding box from [x1, y1, w, h] to [x1, y1, x2, y2] where xy1=top-left, xy2=bottom-right.

    Args:
        x (np.ndarray | torch.Tensor): Input bounding box coordinates.

    Returns:
        (np.ndarray | torch.Tensor): Bounding box coordinates in xyxy format.
    """
    y = x.clone() if isinstance(x, torch.Tensor) else np.copy(x)
    y[..., 2] = x[..., 2] + x[..., 0]  # x2
    y[..., 3] = x[..., 3] + x[..., 1]  # y2
    return y


def segments2boxes(segments):
    """Convert segment coordinates to bounding box labels in xywh format.

    Args:
        segments (list[np.ndarray]): List of segments, each an (N, 2) array of [x, y] points.

    Returns:
        (np.ndarray): Bounding box coordinates in xywh format with shape (N, 4).
    """
    boxes = []
    for s in segments:
        x, y = s.T  # segment xy
        boxes.append([x.min(), y.min(), x.max(), y.max()])  # xyxy
    return xyxy2xywh(np.array(boxes).reshape(-1, 4))  # xywh


def resample_segments(segments, n: int = 1000):
    """Resample closed segments to n points each using linear interpolation.

    Args:
        segments (list): List of (N, 2) arrays where N is the number of points in each segment, modified in place.
        n (int): Number of points to resample each segment to.

    Returns:
        (list): Resampled segments, each an (n, 2) array (segments already of length n are left unchanged).
    """
    for i, s in enumerate(segments):
        if len(s) == n:
            continue
        s = np.concatenate((s, s[0:1, :]), axis=0)
        x = np.linspace(0, len(s) - 1, n - len(s) if len(s) < n else n)
        xp = np.arange(len(s))
        x = np.insert(x, np.searchsorted(x, xp), xp) if len(s) < n else x
        segments[i] = (
            np.concatenate([np.interp(x, xp, s[:, i]) for i in range(2)], dtype=np.float32).reshape(2, -1).T
        )  # segment xy
    return segments


def crop_mask(masks: torch.Tensor, boxes: torch.Tensor) -> torch.Tensor:
    """Crop masks to bounding box regions, zeroing pixels outside each box in place.

    Args:
        masks (torch.Tensor): Masks with shape (N, H, W).
        boxes (torch.Tensor): Bounding box coordinates with shape (N, 4) in xyxy pixel format.

    Returns:
        (torch.Tensor): Cropped masks with shape (N, H, W).
    """
    if boxes.device != masks.device:
        boxes = boxes.to(masks.device)
    _, h, w = masks.shape
    x1, y1, x2, y2 = torch.chunk(boxes[:, :, None], 4, 1)  # each shape (n,1,1)
    r = torch.arange(w, device=masks.device, dtype=x1.dtype)[None, None, :]  # columns (1,1,w)
    c = torch.arange(h, device=masks.device, dtype=x1.dtype)[None, :, None]  # rows (1,h,1)
    # Apply the column and row masks separately and in place: the box region is separable, so this avoids ever
    # materializing the full (n, h, w) boolean grid the combined product would build, and has no per-mask Python loop.
    masks *= (r >= x1) * (r < x2)  # zero columns outside the box
    masks *= (c >= y1) * (c < y2)  # zero rows outside the box
    return masks


def process_mask(protos, masks_in, bboxes, shape, upsample: bool = False):
    """Generate binary instance masks from mask prototypes and coefficients, cropped to their bounding boxes.

    Args:
        protos (torch.Tensor): Mask prototypes with shape (mask_dim, mask_h, mask_w).
        masks_in (torch.Tensor): Mask coefficients with shape (N, mask_dim) where N is number of masks after NMS.
        bboxes (torch.Tensor): Bounding boxes in xyxy format at input image scale with shape (N, 4).
        shape (tuple): Input image size as (height, width).
        upsample (bool): Whether to upsample masks to the input image size.

    Returns:
        (torch.Tensor): A binary uint8 mask tensor of shape [n, h, w], where n is the number of masks after NMS. When
            upsample=True h and w match the input image size; otherwise they are the prototype mask resolution.
    """
    c, mh, mw = protos.shape  # CHW
    if masks_in.shape[0] == 0:  # no detections: F.interpolate below rejects an empty (N=0) batch
        return torch.zeros((0, *(shape if upsample else (mh, mw))), dtype=torch.uint8, device=masks_in.device)
    masks = (masks_in @ protos.float().view(c, -1)).view(-1, mh, mw)  # NHW

    if upsample:
        # Upsample then crop at image resolution; cropping first smears the bilinear edge outside the bbox (#24272)
        masks = F.interpolate(masks[None], shape, mode="bilinear")[0]  # NHW
    else:
        width_ratio = mw / shape[1]
        height_ratio = mh / shape[0]
        ratios = torch.tensor([[width_ratio, height_ratio, width_ratio, height_ratio]], device=bboxes.device)
        bboxes = bboxes * ratios  # scale boxes to prototype resolution
    # Binarize before cropping so crop_mask runs on uint8 instead of float32, as in process_mask_native
    return crop_mask(masks.gt_(0.0).byte(), bboxes)


def process_mask_native(protos, masks_in, bboxes, shape):
    """Generate binary instance masks upsampled to the target image shape with native (letterbox-aware) scaling.

    Args:
        protos (torch.Tensor): Mask prototypes with shape (mask_dim, mask_h, mask_w).
        masks_in (torch.Tensor): Mask coefficients with shape (N, mask_dim) where N is number of masks after NMS.
        bboxes (torch.Tensor): Bounding boxes in xyxy format at the target image scale with shape (N, 4).
        shape (tuple): Target image size as (height, width), typically the original image size.

    Returns:
        (torch.Tensor): Binary uint8 mask tensor with shape (N, H, W), where (H, W) is shape.
    """
    c, mh, mw = protos.shape  # CHW
    h, w = shape
    if masks_in.shape[0] == 0:  # no detections: return a well-formed empty mask stack
        return torch.zeros((0, h, w), dtype=torch.uint8, device=masks_in.device)
    coeffs = masks_in @ protos.float().view(c, -1)  # (N, mh*mw) prototype-resolution mask logits
    # Upsampling all N masks at once allocates an N*H*W float intermediate (~9 GB on a large image with many
    # detections), which OOMs the worker. Upsample in chunks bounded by a pixel budget, thresholding each chunk to
    # uint8 immediately so the float intermediate stays small, then crop the assembled uint8 stack.
    step = max(1, 32_000_000 // (h * w))
    masks = [
        scale_masks(coeffs[i : i + step].view(-1, mh, mw)[None], shape)[0].gt_(0.0).byte()
        for i in range(0, coeffs.shape[0], step)
    ]
    return crop_mask(torch.cat(masks), bboxes)


def scale_masks(
    masks: torch.Tensor,
    shape: tuple[int, int],
    ratio_pad: tuple[tuple[float, float], tuple[float, float]] | None = None,
    padding: bool = True,
    mode: str = "bilinear",
) -> torch.Tensor:
    """Rescale segment masks to target shape.

    Args:
        masks (torch.Tensor): Masks with shape (N, C, H, W).
        shape (tuple[int, int]): Target height and width as (height, width).
        ratio_pad (tuple, optional): Ratio and padding values as ((ratio_h, ratio_w), (pad_w, pad_h)), the letterbox
            gains and its top-left padding.
        padding (bool): Whether masks are based on YOLO-style augmented images with padding.
        mode (str): Interpolation mode, e.g. 'bilinear' for logits or 'nearest' for integer class maps.

    Returns:
        (torch.Tensor): Rescaled float masks with shape (N, C, height, width), or the input unchanged if it already
            matches shape.
    """
    im1_h, im1_w = masks.shape[2:]
    im0_h, im0_w = shape[:2]
    if im1_h == im0_h and im1_w == im0_w:
        return masks
    if masks.shape[1] == 0:  # empty mask stack: F.interpolate rejects a 0-length channel dim
        return masks.new_zeros((*masks.shape[:2], im0_h, im0_w), dtype=torch.float32)

    if ratio_pad is None:  # calculate from im0_shape
        gain_h = gain_w = min(im1_h / im0_h, im1_w / im0_w)  # gain  = old / new
        pad_w, pad_h = (im1_w - round(im0_w * gain_w)) / 2, (im1_h - round(im0_h * gain_h)) / 2  # wh padding
    else:
        (gain_h, gain_w), (pad_w, pad_h) = ratio_pad
    top, left = (round(pad_h - 0.1), round(pad_w - 0.1)) if padding else (0, 0)
    bottom, right = top + round(im0_h * gain_h), left + round(im0_w * gain_w)  # content end, odd pads extra at end
    return F.interpolate(masks[..., top:bottom, left:right].float(), shape, mode=mode)  # NCHW masks


def scale_coords(img1_shape, coords, img0_shape, ratio_pad=None, normalize: bool = False, padding: bool = True):
    """Rescale segment coordinates from img1_shape to img0_shape.

    Args:
        img1_shape (tuple): Source image shape as HWC or HW (supports both).
        coords (torch.Tensor): Coordinates to scale with shape (..., C), C >= 2, with x and y in the first two channels
            of the last dimension, e.g. (N, 2) segments or (N, K, 3) keypoints. Modified in place.
        img0_shape (tuple): Target image shape as HWC or HW (supports both).
        ratio_pad (tuple, optional): Ratio and padding values as ((ratio_h, ratio_w), (pad_w, pad_h)).
        normalize (bool): Whether to normalize coordinates to range [0, 1].
        padding (bool): Whether coordinates are based on YOLO-style augmented images with padding.

    Returns:
        (torch.Tensor): Scaled coordinates, clipped to img0_shape.
    """
    img0_h, img0_w = img0_shape[:2]  # supports both HWC or HW shapes
    if ratio_pad is None:  # calculate from img0_shape
        img1_h, img1_w = img1_shape[:2]  # supports both HWC or HW shapes
        gain = min(img1_h / img0_h, img1_w / img0_w)  # gain  = old / new
        new_h, new_w = round(img0_h * gain), round(img0_w * gain)  # LetterBox rounds each side
        gain_y, gain_x = new_h / img0_h, new_w / img0_w
        pad = round((img1_w - new_w) / 2 - 0.1), round((img1_h - new_h) / 2 - 0.1)
    else:
        gain_y, gain_x = ratio_pad[0]
        pad = ratio_pad[1]

    if padding:
        coords[..., 0] -= pad[0]  # x padding
        coords[..., 1] -= pad[1]  # y padding
    coords[..., 0] /= gain_x
    coords[..., 1] /= gain_y
    coords = clip_coords(coords, img0_shape)
    if normalize:
        coords[..., 0] /= img0_w  # width
        coords[..., 1] /= img0_h  # height
    return coords


def regularize_rboxes(rboxes):
    """Regularize rotated bounding boxes to range [0, pi/2).

    Args:
        rboxes (torch.Tensor): Input rotated boxes with shape (..., 5) in xywhr format.

    Returns:
        (torch.Tensor): Regularized rotated boxes with the same shape, with w and h swapped where needed.
    """
    x, y, w, h, t = rboxes.unbind(dim=-1)
    # Swap edge if t >= pi/2 while not being symmetrically opposite
    swap = t % math.pi >= math.pi / 2
    w_ = torch.where(swap, h, w)
    h_ = torch.where(swap, w, h)
    t = t % (math.pi / 2)
    return torch.stack([x, y, w_, h_, t], dim=-1)  # regularized boxes


def masks2segments(masks: np.ndarray | torch.Tensor, strategy: str = "all") -> list[np.ndarray]:
    """Convert masks to segments using contour detection.

    Args:
        masks (np.ndarray | torch.Tensor): Binary masks with shape (N, H, W).
        strategy (str): Segmentation strategy, either 'all' to merge all contours or 'largest' to keep the contour with
            the most points.

    Returns:
        (list[np.ndarray]): List of (K, 2) float32 segment point arrays, one per mask ((0, 2) if no contour).
    """
    from ultralytics.data.converter import merge_multi_segment

    masks = masks.astype("uint8") if isinstance(masks, np.ndarray) else masks.byte().cpu().numpy()
    segments = []
    for x in np.ascontiguousarray(masks):
        c = cv2.findContours(x, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)[0]
        if c:
            if strategy == "all":  # merge and concatenate all segments
                c = (
                    np.concatenate(merge_multi_segment([x.reshape(-1, 2) for x in c]))
                    if len(c) > 1
                    else c[0].reshape(-1, 2)
                )
            elif strategy == "largest":  # select largest segment
                c = np.array(c[np.array([len(x) for x in c]).argmax()]).reshape(-1, 2)
        else:
            c = np.zeros((0, 2))  # no segments found
        segments.append(c.astype("float32"))
    return segments


def convert_torch2numpy_batch(batch: torch.Tensor) -> np.ndarray:
    """Convert a batch of FP32 torch tensors to NumPy uint8 arrays, changing from BCHW to BHWC layout.

    Args:
        batch (torch.Tensor): Input tensor batch with shape (Batch, Channels, Height, Width) and dtype torch.float32.

    Returns:
        (np.ndarray): Output NumPy array batch with shape (Batch, Height, Width, Channels) and dtype uint8.
    """
    return (batch.permute(0, 2, 3, 1).contiguous() * 255).clamp(0, 255).byte().cpu().numpy()


def clean_str(s):
    """Clean a string by replacing special characters with '_' character.

    Args:
        s (str): A string needing special characters replaced.

    Returns:
        (str): A string with special characters replaced by an underscore _.
    """
    return re.sub(pattern="[|@#!¡·$€%&()=?¿^*;:,¨`><+]", repl="_", string=s)


def empty_like(x):
    """Return an empty torch.Tensor or np.ndarray with the same shape and dtype as the input."""
    return torch.empty_like(x, dtype=x.dtype) if isinstance(x, torch.Tensor) else np.empty_like(x, dtype=x.dtype)


_assignment_solver = None  # resolved once on first call: SciPy's solver if installed, else the NumPy fallback


def linear_sum_assignment(cost_matrix):
    """Solve the rectangular linear sum assignment problem (minimum-cost one-to-one matching).

    Uses `scipy.optimize.linear_sum_assignment` when SciPy is installed (faster compiled C++ solver), and otherwise
    falls back to an equivalent pure-NumPy implementation of the same modified Jonker-Volgenant shortest augmenting path
    algorithm (Crouse 2016). This keeps SciPy out of Ultralytics' required dependencies while preserving its speed when
    present. SciPy is imported lazily so it never slows `import ultralytics`. For a rectangular matrix only min(rows,
    columns) entries are matched.

    The NumPy fallback supports `+inf` as a forbidden assignment and raises `ValueError("cost matrix is infeasible")`
    when no assignment exists; callers must sanitize `NaN` and `-inf`. The two backends may return a different
    equal-cost assignment under exact ties, but the total cost is identical.

    The NumPy fallback is validated against SciPy with exact optimal-cost parity across ~6.9k randomized cases (every
    shape including empty/tall/wide, ties, negatives, IoU- and RT-DETR-style matrices, `maximize` via negation,
    torch-tensor input) plus ~2k independent brute-force global-optimum checks. SciPy's compiled inner loop is faster,
    but at the call-site sizes (smaller dimension = object count) the fallback runs in well under a millisecond:

        cost matrix   NumPy   SciPy
        300 x 20      0.2ms   0.02ms
        300 x 80      0.6ms   0.1ms
        300 x 300     28ms    1.5ms

    Args:
        cost_matrix (np.ndarray | torch.Tensor): Cost matrix with shape (N, M); `+inf` forbids assignments.

    Returns:
        row_ind (np.ndarray): Row indices of the optimal assignment, sorted ascending, with length min(N, M).
        col_ind (np.ndarray): Column indices matched to each row in row_ind.

    Examples:
        >>> cost = np.array([[4, 1, 3], [2, 0, 5], [3, 2, 2]], dtype=float)
        >>> row_ind, col_ind = linear_sum_assignment(cost)
        >>> float(cost[row_ind, col_ind].sum())
        5.0
    """
    global _assignment_solver
    if _assignment_solver is None:  # resolve the backend once, then reuse it on every later call
        try:
            from scipy.optimize import linear_sum_assignment as solver  # faster compiled C++ solver when installed

            _assignment_solver = solver
        except ImportError:
            _assignment_solver = _linear_sum_assignment_numpy
    return _assignment_solver(np.asarray(cost_matrix, dtype=np.float64))


def _linear_sum_assignment_numpy(a):
    """Solve the rectangular linear sum assignment problem with NumPy (Jonker-Volgenant SciPy-free fallback).

    Args:
        a (np.ndarray): Float64 cost matrix of shape (N, M); `+inf` forbids assignments.

    Returns:
        row_ind (np.ndarray): Row indices of the optimal assignment, sorted ascending, with length min(N, M).
        col_ind (np.ndarray): Column indices matched to each row in row_ind.
    """
    n, m = a.shape
    if n == 0 or m == 0:
        return np.empty(0, dtype=np.intp), np.empty(0, dtype=np.intp)
    transposed = n > m
    if transposed:
        a, n, m = a.T, m, n  # ensure rows <= columns
    u, v = np.zeros(n + 1), np.zeros(m + 1)  # row and column dual potentials
    p, way = np.zeros(m + 1, np.intp), np.zeros(m + 1, np.intp)  # column->row matches and path pointers
    for i in range(1, n + 1):
        p[0], j0 = i, 0
        minv, used = np.full(m + 1, np.inf), np.zeros(m + 1, bool)
        while True:  # grow a shortest augmenting path from row i
            used[j0] = True
            i0 = p[j0]
            cur = a[i0 - 1] - u[i0] - v[1:]
            improve = (~used[1:]) & (cur < minv[1:])
            minv[1:][improve], way[1:][improve] = cur[improve], j0
            candidates = np.where(used[1:], np.inf, minv[1:])
            j1 = int(np.argmin(candidates)) + 1
            delta = candidates[j1 - 1]
            if delta == np.inf:
                raise ValueError("cost matrix is infeasible")
            u[p[used]] += delta
            v[used] -= delta
            minv[~used] -= delta
            j0 = j1
            if p[j0] == 0:
                break
        while j0:  # augment along the path
            p[j0] = p[way[j0]]
            j0 = way[j0]
    cols = np.nonzero(p[1:])[0]
    rows = p[1:][cols] - 1
    row_ind, col_ind = (cols, rows) if transposed else (rows, cols)
    order = np.argsort(row_ind, kind="stable")  # match scipy's row-sorted output
    return row_ind[order].astype(np.intp), col_ind[order].astype(np.intp)
