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

from __future__ import annotations

import os
import random
import subprocess
import time
import zipfile
from pathlib import Path
from tarfile import is_tarfile
from typing import Any
from uuid import uuid4

import cv2
import numpy as np
from PIL import Image, ImageOps

from ultralytics.nn.autobackend import check_class_names
from ultralytics.utils import (
    ASSETS_URL,
    DATASETS_DIR,
    LOGGER,
    ROOT,
    SETTINGS_FILE,
    YAML,
    clean_url,
    colorstr,
    emojis,
    is_dir_writeable,
)
from ultralytics.utils.checks import check_file, check_font, is_ascii, normalize_platform_uri
from ultralytics.utils.downloads import download, safe_download
from ultralytics.utils.ops import segments2boxes
from ultralytics.utils.patches import imread

HELP_URL = "See https://docs.ultralytics.com/datasets for dataset formatting guidance."
IMG_FORMATS = {
    "avif",
    "bmp",
    "dng",
    "heic",
    "heif",
    "jp2",
    "jpeg",
    "jpg",
    "mpo",
    "png",
    "tif",
    "tiff",
    "webp",
}
VID_FORMATS = {"asf", "avi", "gif", "m4v", "mkv", "mov", "mp4", "mpeg", "mpg", "ts", "wmv", "webm"}  # videos
FORMATS_HELP_MSG = f"Supported formats are:\nimages: {IMG_FORMATS}\nvideos: {VID_FORMATS}"
DATASET_KEY_TYPES = {  # dataset YAML keys and their permitted types
    "path": (str,),
    "train": (str, list),
    "val": (str, list),
    "test": (str, list),
    "names": (list, dict),
    "kpt_shape": (list,),
    "flip_idx": (list,),
}

DEPTH_PNG_SCALE = 1000  # uint16 millimeters by default; zero is invalid


def save_depth_png(path: str | Path, depth: np.ndarray, scale: float = DEPTH_PNG_SCALE) -> None:
    """Save metric depth as a scaled uint16 PNG with zero reserved for invalid pixels.

    Args:
        path (str | Path): Output PNG file path.
        depth (np.ndarray): Metric depth map in meters, 2D after squeezing. Non-finite and non-positive values are saved
            as 0 (invalid).
        scale (float, optional): Multiplier applied to depth in meters before rounding to uint16, e.g. 1000 for
            millimeters.

    Raises:
        ValueError: If scale is not a positive finite number, depth is not 2D, or scaled depth exceeds the uint16 range.
        OSError: If the PNG cannot be written.
    """
    if not isinstance(scale, (int, float)) or isinstance(scale, bool) or not np.isfinite(scale) or scale <= 0:
        raise ValueError("Depth scale must be a positive finite number")
    depth = np.asarray(depth, dtype=np.float32).squeeze()
    if depth.ndim != 2:
        raise ValueError(f"Depth map must be 2D, got shape {depth.shape}")
    valid = np.isfinite(depth) & (depth > 0)
    encoded = np.zeros(depth.shape, dtype=np.uint16)
    if valid.any():
        scaled = np.rint(depth[valid] * scale)
        if scaled.max() > np.iinfo(np.uint16).max:
            raise ValueError(
                f"Depth map exceeds the {np.iinfo(np.uint16).max / scale:g} meter PNG limit at scale={scale:g}. "
                "Pass a lower scale, e.g. 256, and set the same 'depth_scale' in the dataset YAML."
            )
        encoded[valid] = np.maximum(scaled, 1).astype(np.uint16)
    if not cv2.imwrite(str(path), encoded):
        raise OSError(f"Failed to save depth map to {path}")


def load_depth(path: str | Path, scale: float = DEPTH_PNG_SCALE) -> np.ndarray:
    """Load metric depth from a scaled uint16 PNG or floating-point meter NPY.

    Args:
        path (str | Path): Path to a *.png depth map (uint16 scaled by `scale`) or *.npy depth map (float, meters).
        scale (float, optional): Divisor applied to PNG values to convert them to meters. Ignored for NPY files.

    Returns:
        (np.ndarray): Float32 depth map in meters with shape (H, W), where 0 marks invalid pixels.

    Raises:
        ValueError: If the depth file has an unsupported shape, dtype, or format, or scale is not a positive finite
            number.
    """
    path = Path(path)
    if path.suffix.lower() == ".npy":
        depth = np.load(path, allow_pickle=False)
        if depth.ndim != 2 or depth.dtype.kind != "f":
            raise ValueError(f"Depth map {path} must be a 2D floating-point NPY array")
        return np.nan_to_num(depth.astype(np.float32, copy=False), copy=False, nan=0.0, posinf=0.0, neginf=0.0)
    if not isinstance(scale, (int, float)) or isinstance(scale, bool) or not np.isfinite(scale) or scale <= 0:
        raise ValueError("Depth scale must be a positive finite number")
    with Image.open(path) as image:
        if image.format != "PNG" or image.mode not in {"I", "I;16"}:
            raise ValueError(f"Depth PNG {path} must be a 2D uint16 map")
        encoded = np.asarray(image)
    depth = encoded.astype(np.float32)
    depth /= scale
    return depth


def img2label_paths(img_paths: list[str | Path], label_dir: str = "labels", suffix: str = ".txt") -> list[str]:
    """Convert image paths to label paths by replacing the last 'images' directory and the file extension.

    Args:
        img_paths (list[str | Path]): List of image file paths.
        label_dir (str, optional): Directory name that replaces the last '/images/' path component.
        suffix (str, optional): File extension that replaces the image extension.

    Returns:
        (list[str]): List of label file paths.
    """
    sa, sb = f"{os.sep}images{os.sep}", f"{os.sep}{label_dir}{os.sep}"  # /images/, /labels/ substrings
    return [sb.join(os.fspath(x).rsplit(sa, 1)).rsplit(".", 1)[0] + f"{suffix}" for x in img_paths]


def check_file_speeds(
    files: list[str | Path], threshold_ms: float = 10, threshold_mb: float = 50, max_files: int = 5, prefix: str = ""
):
    """Check dataset file access speed and provide performance feedback.

    This function tests the access speed of dataset files by measuring ping (stat call) time and read speed. It samples
    up to `max_files` files from the provided list and warns if access times exceed the threshold.

    Args:
        files (list[str | Path]): List of file paths to check for access speed.
        threshold_ms (float, optional): Threshold in milliseconds for ping time warnings.
        threshold_mb (float, optional): Threshold in megabytes per second for read speed warnings.
        max_files (int, optional): The maximum number of files to check.
        prefix (str, optional): Prefix string to add to log messages.

    Examples:
        >>> from pathlib import Path
        >>> image_files = list(Path("dataset/images").glob("*.jpg"))
        >>> check_file_speeds(image_files, threshold_ms=15)
    """
    if not files:
        LOGGER.warning(f"{prefix}Image speed checks: No files to check")
        return

    # Sample up to max_files files
    files = random.sample(files, min(max_files, len(files)))

    # Test ping (stat time)
    ping_times = []
    file_sizes = []
    read_speeds = []

    for f in files:
        try:
            # Measure ping (stat call)
            start = time.perf_counter()
            file_size = os.stat(f).st_size
            ping_times.append((time.perf_counter() - start) * 1000)  # ms
            file_sizes.append(file_size)

            # Measure read speed
            start = time.perf_counter()
            with open(f, "rb") as file_obj:
                _ = file_obj.read()
            read_time = time.perf_counter() - start
            if read_time > 0:  # Avoid division by zero
                read_speeds.append(file_size / (1 << 20) / read_time)  # MB/s
        except Exception:
            pass

    if not ping_times:
        LOGGER.warning(f"{prefix}Image speed checks: failed to access files")
        return

    # Calculate stats with uncertainties
    avg_ping = np.mean(ping_times)
    std_ping = np.std(ping_times, ddof=1) if len(ping_times) > 1 else 0
    size_msg = f", size: {np.mean(file_sizes) / (1 << 10):.1f} KB"
    ping_msg = f"ping: {avg_ping:.1f}±{std_ping:.1f} ms"

    if read_speeds:
        avg_speed = np.mean(read_speeds)
        std_speed = np.std(read_speeds, ddof=1) if len(read_speeds) > 1 else 0
        speed_msg = f", read: {avg_speed:.1f}±{std_speed:.1f} MB/s"
    else:
        avg_speed = float("inf")
        speed_msg = ""

    # MB/s is open() latency-bound for tiny files (0.2 KB mnist160 PNGs read ~15 MB/s on local NVMe), so skip it there
    if avg_ping < threshold_ms and (avg_speed > threshold_mb or np.mean(file_sizes) < 1 << 14):
        LOGGER.info(f"{prefix}Fast image access ✅ ({ping_msg}{speed_msg}{size_msg})")
    else:
        LOGGER.warning(
            f"{prefix}Slow image access detected ({ping_msg}{speed_msg}{size_msg}). "
            f"Use local storage instead of remote/mounted storage for better performance. "
            f"See https://docs.ultralytics.com/guides/model-training-tips"
        )


def get_hash(paths: list[str]) -> str:
    """Return a hash of paths and their file sizes and modification times."""
    h = __import__("hashlib").sha256()
    for p in paths:
        h.update(p.encode())
        h.update(b"\0")
        try:
            stat = os.stat(p)
        except OSError:
            h.update(b"\0")
            continue
        h.update(f"{stat.st_size}:{stat.st_mtime_ns}".encode())
        h.update(b"\0")
    return h.hexdigest()


def exif_size(img: Image.Image) -> tuple[int, int]:
    """Return exif-corrected PIL size."""
    s = img.size  # (width, height)
    try:
        exif = img.tag_v2 if img.format == "TIFF" else img.getexif()  # TIFF tags stay readable after verify()
        if exif.get(274) in {5, 6, 7, 8}:  # swap w and h; WebP/TIFF vary by cv2/Pillow version, so decode those
            s = s[::-1] if img.format in {"JPEG", "MPO", "PNG", "AVIF"} else imread(img.filename).shape[1::-1]
    except Exception:
        pass
    return s


def check_image(im_file: str) -> tuple[str, tuple[int, int]]:
    """Verify an image file for integrity and correct corrupt JPEGs if found.

    Args:
        im_file (str): Path to the image file to check.

    Returns:
        (str): A message describing any corrective action taken, or an empty string if the image is valid.
        (tuple[int, int]): Image shape as (height, width) in pixels.

    Raises:
        AssertionError: If the image size is less than 10 pixels in any dimension or the format is invalid.
    """
    msg = ""
    im = Image.open(im_file)
    im.verify()  # PIL verify
    shape = exif_size(im)  # image size
    shape = (shape[1], shape[0])  # hw
    assert (shape[0] > 9) & (shape[1] > 9), f"image size {shape} <10 pixels"
    assert im.format.lower() in IMG_FORMATS | {"jpeg2000"}, f"Invalid image format {im.format}. {FORMATS_HELP_MSG}"
    if im.format.lower() in {"jpg", "jpeg"}:
        with open(im_file, "rb") as f:
            f.seek(-2, 2)
            corrupt = f.read() != b"\xff\xd9"
        if corrupt:  # write a new file and swap it in: the image may be a hard link shared with other versions
            _replace_image(im_file, lambda tmp: _exif_jpeg(im_file).save(tmp, "JPEG", subsampling=0, quality=100))
            msg = f"{im_file}: corrupt JPEG restored and saved"
    return msg, shape


def _exif_jpeg(im_file: str | Path) -> Image.Image:
    """Load an image with its EXIF orientation applied, closing the source file."""
    with Image.open(im_file) as im:
        return ImageOps.exif_transpose(im)


def _replace_image(im_file: str | Path, write) -> None:
    """Atomically replace an image with what `write(tmp)` saves, so hard links to the original are never modified."""
    im_file = Path(im_file)
    tmp = im_file.with_name(f".{im_file.stem}.{uuid4().hex}{im_file.suffix}")
    try:
        write(str(tmp))
        os.replace(tmp, im_file)
    finally:
        tmp.unlink(missing_ok=True)


def verify_image(args: tuple) -> tuple:
    """Verify one image for classification datasets.

    Args:
        args (tuple): Tuple of ((im_file, cls), prefix).

    Returns:
        (tuple): Tuple of ((im_file, cls), nf, nc, msg), where nf and nc are 1 if the image was found valid or corrupt
            respectively, and msg is a log message.
    """
    (im_file, cls), prefix = args
    # Number (found, corrupt), message
    nf, nc, msg = 0, 0, ""
    try:
        msg = check_image(im_file)[0]
        msg = f"{prefix}{msg}" if msg else ""
        nf = 1
    except Exception as e:
        nc = 1
        msg = f"{prefix}{im_file}: ignoring corrupt image/label: {e}"
    return (im_file, cls), nf, nc, msg


def verify_image_depth(args: tuple) -> tuple:
    """Verify that an image and its paired depth map exist and are readable.

    Args:
        args (tuple): Tuple of (im_file, depth_file, prefix, scale).

    Returns:
        (tuple): Tuple of (im_file, shape, nf, nm, nc, msg), where im_file and shape (H, W) are None for rejected
            samples, nf, nm, and nc are found, missing, and corrupt counts, and msg is a log message.
    """
    im_file, depth_file, prefix, scale = args
    # Number (found, missing, corrupt), message
    nf, nm, nc, msg = 0, 0, 0, ""
    try:
        msg, shape = check_image(im_file)
        msg = f"{prefix}{msg}" if msg else ""
        if not os.path.isfile(depth_file):
            nm = 1
            msg = f"{prefix}{im_file}: ignoring image with missing depth map {depth_file}"
            return None, None, nf, nm, nc, msg
        if Path(depth_file).suffix.lower() == ".npy":
            depth = np.load(depth_file, mmap_mode="r", allow_pickle=False)
            assert depth.ndim == 2 and depth.dtype.kind == "f", "depth NPY must be 2D and floating-point"
            depth_shape = depth.shape
        else:
            assert (
                isinstance(scale, (int, float)) and not isinstance(scale, bool) and np.isfinite(scale) and scale > 0
            ), "depth_scale must be a positive finite number"
            with Image.open(depth_file) as depth:
                assert depth.format == "PNG" and depth.mode in {"I", "I;16"}, (
                    f"depth map {depth_file} must be an integer grayscale PNG"
                )
                depth_shape = (depth.height, depth.width)
                depth.verify()
        assert abs(np.log((depth_shape[1] / depth_shape[0]) / (shape[1] / shape[0]))) <= 0.02, (
            f"depth map shape {depth_shape} does not match image shape {shape}"
        )
        nf = 1
        return im_file, shape, nf, nm, nc, msg
    except Exception as e:
        nc = 1
        msg = f"{prefix}{im_file}: ignoring corrupt image/depth: {e}"
    return None, None, nf, nm, nc, msg


def verify_image_mask(args: tuple) -> tuple:
    """Verify that an image and its semantic mask exist, are readable, match in shape, and hold valid class ids.

    Args:
        args (tuple): Tuple of (im_file, mask_file, prefix, invalid). If mask_file is missing, masks with the same stem
            and another image extension are tried. invalid is a 256-entry uint8 lookup table that is nonzero for raw
            mask ids that map to neither a dataset class nor the 255 ignore label.

    Returns:
        (tuple): Tuple of (im_file, mask_file, shape, mode, nm, nf, nc, msg), where the first four are None for rejected
            samples, mode is the mask's PIL image mode, nm, nf, and nc are missing, found, and corrupt counts, and msg
            is a log message.
    """
    im_file, mask_file, prefix, invalid = args
    # Number (found, missing, corrupt), message
    nf, nm, nc, msg = 0, 0, 0, ""
    try:
        msg, shape = check_image(im_file)
        msg = f"{prefix}{msg}" if msg else ""
        if not os.path.isfile(mask_file):
            for ext in IMG_FORMATS:  # check other suffixes
                alt_mask_file = mask_file.rsplit(".", 1)[0] + f".{ext}"
                if os.path.isfile(alt_mask_file):
                    mask_file = alt_mask_file
                    break
        if os.path.isfile(mask_file):
            with Image.open(mask_file) as im:
                mode = im.mode  # recorded so load_mask reads each mask once and a yaml 'nc' edit never needs a rescan
                if mode == "P":  # colored (VOC-style) palettes hold class ids as indices, gray palettes as gray levels
                    p = np.array(im.getpalette()).reshape(-1, 3)
                    mask = np.asarray(im.convert("L") if (p == p[:, :1]).all() else im)
                else:
                    mask = cv2.imread(mask_file, cv2.IMREAD_ANYDEPTH)  # keeps 16-bit ids
            assert mask is not None, f"mask file {mask_file} is unreadable"
            assert mask.shape[:2] == shape, f"mask size {mask.shape[:2]} does not match image size {shape}"
            assert not invalid[mask].any(), (  # ids above 255 raise IndexError
                f"mask ids {np.unique(mask[invalid[mask] > 0]).tolist()} are not dataset class ids or 255 ignore"
            )
            nf = 1
        else:
            nm = 1
            msg = f"{prefix}{im_file}: ignoring image with missing mask {mask_file}"
            return None, None, None, None, nm, nf, nc, msg
        return im_file, mask_file, shape, mode, nm, nf, nc, msg
    except Exception as e:
        nc = 1
        msg = f"{prefix}{im_file}: ignoring corrupt image/mask: {e}"
    return None, None, None, None, nm, nf, nc, msg


def verify_image_label(args: tuple) -> tuple | list:
    """Verify one image-label pair.

    Args:
        args (tuple): Tuple of (im_file, lb_file, prefix, keypoint, num_cls, nkpt, ndim, single_cls).

    Returns:
        (tuple | list): Tuple of (im_file, lb, shape, segments, keypoints, nm, nf, ne, nc, msg), where lb is an (N, 5)
            array of [cls, x, y, w, h] labels, shape is (H, W), segments is a list of (K, 2) arrays, keypoints is an (N,
            nkpt, 3) array or None, nm, nf, ne, and nc are missing, found, empty, and corrupt counts, and msg is a log
            message. For corrupt samples, a list with the first five items set to None is returned.
    """
    im_file, lb_file, prefix, keypoint, num_cls, nkpt, ndim, single_cls = args
    # Number (missing, found, empty, corrupt), message, segments, keypoints
    nm, nf, ne, nc, msg, segments, keypoints = 0, 0, 0, 0, "", [], None
    try:
        # Verify images
        msg, shape = check_image(im_file)
        msg = f"{prefix}{msg}" if msg else ""

        # Verify labels
        if os.path.isfile(lb_file):
            nf = 1  # label found
            with open(lb_file, encoding="utf-8") as f:
                lb = [x.split() for x in f.read().strip().splitlines() if x.strip()]
                if nkpt and not keypoint:  # pose labels for a box task: keep the box, drop the keypoints
                    lb = [x[:5] if len(x) == 5 + nkpt * ndim else x for x in lb]
                if any(len(x) > 6 for x in lb) and (not keypoint):  # is segment
                    assert not any(len(x) == 5 for x in lb), "labels mix segment and detection rows"
                    classes = np.array([x[0] for x in lb], dtype=np.float32)
                    segments = [np.array(x[1:], dtype=np.float32).reshape(-1, 2) for x in lb]  # (cls, xy1...)
                    lb = np.concatenate((classes.reshape(-1, 1), segments2boxes(segments)), 1)  # (cls, xywh)
                lb = np.array(lb, dtype=np.float32)
            if nl := len(lb):
                if keypoint:
                    assert lb.shape[1] == (5 + nkpt * ndim), f"labels require {(5 + nkpt * ndim)} columns each"
                    points = lb[:, 5:].reshape(-1, ndim)[:, :2]
                else:
                    assert lb.shape[1] == 5, f"labels require 5 columns, {lb.shape[1]} columns detected"
                    points = lb[:, 1:]
                # Coordinate points check with 1% tolerance
                assert points.max() <= 1.01, f"non-normalized or out of bounds coordinates {points[points > 1.01]}"
                assert lb.min() >= -0.01, f"negative class labels or coordinate {lb[lb < -0.01]}"
                assert (lb[:, 0] % 1 == 0).all(), f"non-integer class labels {lb[:, 0][lb[:, 0] % 1 != 0]}"

                # All labels
                max_cls = 0 if single_cls else lb[:, 0].max()  # max class index
                assert max_cls < num_cls, (
                    f"Label class {int(max_cls)} exceeds dataset class count {num_cls}. "
                    f"Possible class labels are 0-{num_cls - 1}"
                )
                _, i = np.unique(lb, axis=0, return_index=True)
                if len(i) < nl and segments:  # distinct polygons can share a class and box
                    rows = np.array([c.tobytes() + s.tobytes() for c, s in zip(lb[:, 0], segments)], dtype=object)
                    _, i = np.unique(rows, return_index=True)
                if len(i) < nl:  # duplicate row check
                    lb = lb[i]  # remove duplicates
                    if segments:
                        segments = [segments[x] for x in i]
                    msg = f"{prefix}{im_file}: {nl - len(i)} duplicate labels removed"
            else:
                ne = 1  # label empty
                lb = np.zeros((0, (5 + nkpt * ndim) if keypoint else 5), dtype=np.float32)
        else:
            nm = 1  # label missing
            lb = np.zeros((0, (5 + nkpt * ndim) if keypoint else 5), dtype=np.float32)
        if keypoint:
            keypoints = lb[:, 5:].reshape(-1, nkpt, ndim)
            if ndim == 2:
                kpt_mask = np.where((keypoints[..., 0] < 0) | (keypoints[..., 1] < 0), 0.0, 1.0).astype(np.float32)
                keypoints = np.concatenate([keypoints, kpt_mask[..., None]], axis=-1)  # (nl, nkpt, 3)
        lb = lb[:, :5]
        return im_file, lb, shape, segments, keypoints, nm, nf, ne, nc, msg
    except Exception as e:
        nc = 1
        msg = f"{prefix}{im_file}: ignoring corrupt image/label: {e}"
        return [None, None, None, None, None, nm, nf, ne, nc, msg]


def visualize_image_annotations(image_path: str, txt_path: str, label_map: dict[int, str]):
    """Visualize YOLO detection annotations (bounding boxes and class labels) on an image.

    This function reads an image and its corresponding YOLO detection label file, then draws bounding boxes around
    detected objects and labels them with their respective class names. The bounding box colors are assigned based on
    the class ID, and the text color is dynamically adjusted for readability, depending on the background color's
    luminance.

    Args:
        image_path (str): Path to the image file to annotate. The file must be readable by PIL.
        txt_path (str): Path to a YOLO detection label file with one `class x_center y_center width height` line per
            object. Segmentation polygon and pose label rows are not supported.
        label_map (dict[int, str]): A dictionary that maps class IDs (integers) to class labels (strings).

    Examples:
        >>> label_map = {0: "cat", 1: "dog", 2: "bird"}  # Should include all annotated classes
        >>> visualize_image_annotations("path/to/image.jpg", "path/to/annotations.txt", label_map)
    """
    import matplotlib.pyplot as plt

    from ultralytics.utils.plotting import colors

    img = np.array(ImageOps.exif_transpose(Image.open(image_path)))  # upright, as dataloaders read it for training
    img_height, img_width = img.shape[:2]
    annotations = []
    with open(txt_path, encoding="utf-8") as file:
        for line in file:
            class_id, x_center, y_center, width, height = map(float, line.split())
            x = (x_center - width / 2) * img_width
            y = (y_center - height / 2) * img_height
            w = width * img_width
            h = height * img_height
            annotations.append((x, y, w, h, int(class_id)))
    _, ax = plt.subplots(1)  # Plot the image and annotations
    for x, y, w, h, label in annotations:
        color = tuple(c / 255 for c in colors(label, False))  # Get and normalize an RGB color for Matplotlib
        rect = plt.Rectangle((x, y), w, h, linewidth=2, edgecolor=color, facecolor="none")  # Create a rectangle
        ax.add_patch(rect)
        luminance = 0.2126 * color[0] + 0.7152 * color[1] + 0.0722 * color[2]  # Formula for luminance
        ax.text(x, y - 5, label_map[label], color="white" if luminance < 0.5 else "black", backgroundcolor=color)
    ax.imshow(img)
    plt.show()


def polygon2mask(
    imgsz: tuple[int, int], polygons: list[np.ndarray], color: int = 1, downsample_ratio: int = 1
) -> np.ndarray:
    """Convert a list of polygons to a binary mask of the specified image size.

    Args:
        imgsz (tuple[int, int]): The size of the image as (height, width).
        polygons (list[np.ndarray]): A list of polygons. Each polygon is a 1D array of coordinates with length M, where
            M % 2 = 0 (alternating x, y values).
        color (int, optional): The color value to fill in the polygons on the mask.
        downsample_ratio (int, optional): Factor by which to downsample the mask.

    Returns:
        (np.ndarray): Mask of shape (H // downsample_ratio, W // downsample_ratio) with the polygons filled with
            `color`.
    """
    mask = np.zeros(imgsz, dtype=np.uint8)
    polygons = np.asarray(polygons, dtype=np.int32)
    polygons = polygons.reshape((polygons.shape[0], -1, 2))
    cv2.fillPoly(mask, polygons, color=color)
    nh, nw = (imgsz[0] // downsample_ratio, imgsz[1] // downsample_ratio)
    # Note: fillPoly first then resize is trying to keep the same loss calculation method when mask-ratio=1
    return cv2.resize(mask, (nw, nh))


def polygons2masks(
    imgsz: tuple[int, int], polygons: list[np.ndarray], color: int, downsample_ratio: int = 1
) -> np.ndarray:
    """Convert a list of polygons to a set of binary masks of the specified image size.

    Args:
        imgsz (tuple[int, int]): The size of the image as (height, width).
        polygons (list[np.ndarray]): A list of polygons. Each polygon is an array of coordinates that can be reshaped to
            (-1, 2) as (x, y) point pairs.
        color (int): The color value to fill in the polygons on the masks.
        downsample_ratio (int, optional): Factor by which to downsample each mask.

    Returns:
        (np.ndarray): Masks of shape (N, H // downsample_ratio, W // downsample_ratio), one per polygon, filled with
            `color`.
    """
    return np.array([polygon2mask(imgsz, [x.reshape(-1)], color, downsample_ratio) for x in polygons])


def polygons2masks_overlap(
    imgsz: tuple[int, int], segments: list[np.ndarray], downsample_ratio: int = 1
) -> tuple[np.ndarray, np.ndarray]:
    """Return a downsampled overlap mask and sorted area indices.

    Args:
        imgsz (tuple[int, int]): The size of the image as (height, width).
        segments (list[np.ndarray]): A list of polygons, each reshapeable to (-1, 2) as (x, y) point pairs.
        downsample_ratio (int, optional): Factor by which to downsample the mask.

    Returns:
        masks (np.ndarray): Mask of shape (H // downsample_ratio, W // downsample_ratio) where 0 is background and i + 1
            marks the i-th instance in area-descending order, so smaller instances are drawn over larger ones.
        index (np.ndarray): Indices that sort the segments by area in descending order.
    """
    masks = np.zeros(
        (imgsz[0] // downsample_ratio, imgsz[1] // downsample_ratio),
        dtype=np.int32 if len(segments) > 255 else np.uint8,
    )
    areas = []
    ms = []
    for segment in segments:
        mask = polygon2mask(
            imgsz,
            [segment.reshape(-1)],
            downsample_ratio=downsample_ratio,
            color=1,
        )
        ms.append(mask.astype(masks.dtype))
        areas.append(mask.sum())
    areas = np.asarray(areas)
    index = np.argsort(-areas)
    ms = np.array(ms)[index]
    # Running max: the old `masks + mask` sum hit 2 * i + 1 and overflowed uint8 past 128 overlapping instances
    for i in range(len(segments)):
        np.maximum(masks, ms[i] * (i + 1), out=masks)
    return masks, index


def find_dataset_yaml(path: Path) -> Path:
    """Find and return the YAML file associated with a Detect, Segment or Pose dataset.

    This function searches for a YAML file at the root level of the provided directory first, and if not found, it
    performs a recursive search. It prefers YAML files that have the same stem as the provided path.

    Args:
        path (Path): The directory path to search for the YAML file.

    Returns:
        (Path): The path of the found YAML file.
    """
    files = list(path.glob("*.yaml")) or list(path.rglob("*.yaml"))  # try root level first and then recursive
    assert files, f"No YAML file found in '{path.resolve()}'"
    if len(files) > 1:
        files = [f for f in files if f.stem == path.stem]  # prefer YAML files that match
    assert len(files) == 1, f"Expected 1 YAML file in '{path.resolve()}', but found {len(files)}.\n{files}"
    return files[0]


def get_split_fraction(fraction: float | list[float | int], split: str) -> float | int:
    """Return a split ratio/count, normalizing boundary values to 0.0 (none) or 1.0 (all).

    Args:
        fraction (float | int | list[float | int]): Dataset fraction (ratio or image count), or a per-split list ordered
            as [train, val, test]. A scalar only applies to the train split; missing list entries default to 1.0.
        split (str): Dataset split name, e.g. 'train', 'val', or 'test'.

    Returns:
        (float | int): Fraction of the split to use as a ratio (float) or image count (int).

    Raises:
        ValueError: If the resolved fraction is 0 for the 'train' or 'val' split.
    """
    if isinstance(fraction, list) and split in (splits := ("train", "val", "test")):
        index = splits.index(split)
        fraction = fraction[index] if index < len(fraction) else 1.0
    elif split != "train":
        fraction = 1.0
    fraction = float(fraction) if fraction in {0, 1} else fraction
    if split in {"train", "val"} and fraction == 0:
        raise ValueError(f"{split} fraction must select at least one image")
    return fraction


def convert_ndjson_to_yolo_if_needed(
    data: str | Path, fraction: float | list[float | int] = 1.0, *, split: str | None = None
) -> str | Path:
    """Convert an NDJSON dataset or Platform dataset URI to YOLO format.

    Args:
        data (str | Path): Dataset path, NDJSON file path or URL, or Ultralytics Platform dataset URI or web URL.
        fraction (float | int | list[float | int], optional): Dataset fraction passed to the NDJSON converter.
        split (str, optional): Dataset split passed to the NDJSON converter.

    Returns:
        (str | Path): Path to the converted dataset (YAML file or directory) for NDJSON inputs, otherwise the normalized
            input data unchanged.
    """
    data = normalize_platform_uri(data)  # accept Platform web URLs (https://platform.ultralytics.com/.../datasets/...)
    data_str = str(data)
    if clean_url(data_str).endswith(".ndjson") or (data_str.startswith("ul://") and "/datasets/" in data_str):
        import asyncio

        from ultralytics.data.converter import convert_ndjson_to_yolo

        return asyncio.run(convert_ndjson_to_yolo(data, fraction=fraction, split=split))
    return data


def check_det_dataset(dataset: str | Path, autodownload: bool = True, split: str = "") -> dict[str, Any]:
    """Download, verify, and/or unzip a dataset if not found locally.

    This function checks the availability of a specified dataset, and if not found, it has the option to download and
    unzip the dataset. It then reads and parses the accompanying YAML data, ensuring key requirements are met and also
    resolves paths related to the dataset.

    Args:
        dataset (str | Path): Path to the dataset or dataset descriptor (like a YAML file).
        autodownload (bool, optional): Whether to automatically download the dataset if not found.
        split (str, optional): Dataset split required by the caller.

    Returns:
        (dict[str, Any]): Parsed dataset information and paths.
    """
    dataset = str(dataset)
    if "://" not in dataset and not Path(dataset).exists() and Path(dataset).suffix not in {".yaml", ".yml"}:
        # allow bare dataset names, e.g. 'coco8' -> 'coco8.yaml', 'DOTAv1.5' -> 'DOTAv1.5.yaml'
        dataset = next((f"{dataset}{x}" for x in (".yaml", ".yml") if check_file(f"{dataset}{x}", hard=False)), dataset)
    file = Path(check_file(dataset))
    if file.is_dir():
        file = find_dataset_yaml(file)

    # Download (optional)
    extract_dir = ""
    if zipfile.is_zipfile(file) or is_tarfile(file):
        new_dir = safe_download(file, dir=DATASETS_DIR, unzip=True, delete=False)
        file = new_dir if new_dir.is_file() else find_dataset_yaml(new_dir)
        extract_dir, autodownload = file.parent, False

    # Read YAML
    data = YAML.load(file, append_filename=True)  # dictionary

    # Checks
    for key, valid_types in DATASET_KEY_TYPES.items():
        if data.get(key) is not None and not isinstance(data[key], valid_types):
            expected = " or ".join(t.__name__ for t in valid_types)
            raise TypeError(f"{dataset} '{key}' must be {expected}, not {type(data[key]).__name__}")

    for k in "train", "val":
        if k not in data:
            if k != "val" or "validation" not in data:
                raise SyntaxError(
                    emojis(f"{dataset} '{k}:' key missing ❌.\n'train' and 'val' are required in all data YAMLs.")
                )
            LOGGER.warning("renaming data YAML 'validation' key to 'val' to match YOLO format.")
            data["val"] = data.pop("validation")  # replace 'validation' key with 'val' key
    if split and not data.get(split):
        raise FileNotFoundError(f"{dataset} '{split}:' images not found ❌")
    # `names` compared to None, not membership: a bare `names:` parses to None and len(None) below
    # raises. `nc` stays membership so a valueless `nc:` still reaches its "must be an integer" error.
    if data.get("names") is None and "nc" not in data:
        raise SyntaxError(emojis(f"{dataset} key missing ❌.\n either 'names' or 'nc' are required in all data YAMLs."))
    if "nc" in data and not isinstance(data["nc"], int):
        try:
            nc = float(data["nc"])  # accept integer-like values, e.g. '10' or 10.0, but not 1.9 or placeholders
            if nc != int(nc):
                raise ValueError
            data["nc"] = int(nc)
        except (TypeError, ValueError):
            raise SyntaxError(emojis(f"{dataset} 'nc: {data['nc']}' must be an integer ❌."))
    if data.get("names") is not None and data.get("nc") is not None and len(data["names"]) != data["nc"]:
        raise SyntaxError(emojis(f"{dataset} 'names' length {len(data['names'])} and 'nc: {data['nc']}' must match."))
    if data.get("names") is None:
        data["names"] = [f"class_{i}" for i in range(data["nc"])]
    else:
        data["nc"] = len(data["names"])

    data["names"] = check_class_names(data["names"])
    data["channels"] = data.get("channels", 3)  # get image channels, default to 3

    # Resolve paths
    path = Path(extract_dir or data.get("path") or Path(data.get("yaml_file", "")).parent)  # dataset root
    if not path.exists() and not path.is_absolute():
        path = (DATASETS_DIR / path).resolve()  # path relative to DATASETS_DIR

    # Set paths
    data["path"] = path  # download scripts
    for k in "train", "val", "test", "minival":
        if data.get(k):  # prepend path
            if isinstance(data[k], str):
                x = (path / data[k]).resolve()
                if not x.exists() and data[k].startswith("../"):
                    x = (path / data[k][3:]).resolve()
                data[k] = str(x)
            else:
                data[k] = [str((path / x).resolve()) for x in data[k]]

    # Parse YAML
    val, s = (data.get(x) for x in (split or "val", "download"))
    if val:
        val = [Path(x).resolve() for x in (val if isinstance(val, list) else [val])]  # val path
        if not all(x.exists() for x in val):
            name = clean_url(dataset)  # dataset name with URL auth stripped
            LOGGER.info("")
            m = f"Dataset '{name}' images not found, missing path '{next(x for x in val if not x.exists())}'"
            if s and autodownload:
                LOGGER.warning(m)
            else:
                m += f"\nNote dataset download directory is '{DATASETS_DIR}'. You can update this in '{SETTINGS_FILE}'"
                raise FileNotFoundError(m)
            t = time.time()
            r = None  # success
            if s.startswith("http") and s.endswith(
                (".zip", ".tar", ".gz", ".tgz", ".xz", ".bz2", ".txz", ".tbz2")
            ):  # URL
                safe_download(url=s, dir=DATASETS_DIR, delete=True)
            elif s.startswith("bash "):  # bash script
                LOGGER.info(f"Running {s} ...")
                subprocess.run(s.split(), check=True)
            else:  # python script
                exec(s, {"yaml": data})  # noqa: S102
            dt = f"({round(time.time() - t, 1)}s)"
            s = f"success ✅ {dt}, saved to {colorstr('bold', DATASETS_DIR)}" if r in {0, None} else f"failure {dt} ❌"
            LOGGER.info(f"Dataset download {s}\n")
    if data.get("masks_dir") is None and (path / "masks").is_dir():  # after download so scripts can create it
        data["masks_dir"] = "masks"  # PNG semantic masks in the default folder select SemanticDataset
    check_font("Arial.ttf" if is_ascii(data["names"]) else "Arial.Unicode.ttf")  # download fonts

    return data  # dictionary


def check_cls_dataset(dataset: str | Path, split: str = "") -> dict[str, Any]:
    """Check a classification dataset such as Imagenet.

    This function accepts a `dataset` name and attempts to retrieve the corresponding dataset information. If the
    dataset is not found locally, it attempts to download the dataset from the internet and save it locally.

    Args:
        dataset (str | Path): The dataset name, local directory path, archive file, or archive URL.
        split (str, optional): The split of the dataset. Either 'train', 'val', 'test', or ''.

    Returns:
        (dict[str, Any]): A dictionary containing the following keys:

            - 'train' (Path): The directory path containing the training set of the dataset.
            - 'val' (Path | None): The directory path containing the validation set of the dataset.
            - 'test' (Path | None): The directory path containing the test set of the dataset.
            - 'nc' (int): The number of classes in the dataset.
            - 'names' (dict[int, str]): A dictionary of class names in the dataset.
            - 'channels' (int): The number of image channels, always 3.
    """
    if split and split not in {"train", "val", "test"}:
        raise ValueError(f"Invalid classification dataset split '{split}'. Use 'train', 'val', or 'test'.")

    # Download (optional if dataset=https://file.zip is passed directly)
    if str(dataset).startswith(("http:/", "https:/")):
        dataset = safe_download(dataset, dir=DATASETS_DIR, unzip=True, delete=False)
    elif str(dataset).endswith((".zip", ".tar", ".gz", ".tgz", ".xz", ".bz2", ".txz", ".tbz2")):
        file = check_file(dataset)
        dataset = safe_download(file, dir=DATASETS_DIR, unzip=True, delete=False)

    dataset = Path(dataset)
    data_dir = (dataset if dataset.is_dir() else (DATASETS_DIR / dataset)).resolve()
    if not data_dir.is_dir():
        if data_dir.suffix != "":
            raise ValueError(
                f'Classification datasets must be a directory (data="path/to/dir") not a file (data="{dataset}"), '
                "See https://docs.ultralytics.com/datasets/classify"
            )
        LOGGER.info("")
        LOGGER.warning(f"Dataset not found, missing path {data_dir}, attempting download...")
        t = time.time()
        if str(dataset) == "imagenet":
            subprocess.run(["bash", str(ROOT / "data/scripts/get_imagenet.sh")], check=True)
        else:
            download(f"{ASSETS_URL}/{dataset}.zip", dir=data_dir.parent)
        LOGGER.info(f"Dataset download success ✅ ({time.time() - t:.1f}s), saved to {colorstr('bold', data_dir)}\n")
    train_set = data_dir / "train"
    if not train_set.is_dir():
        LOGGER.warning(f"Dataset 'split=train' not found at {train_set}")
        if image_files := [f for f in data_dir.rglob("*.*") if f.suffix[1:].lower() in IMG_FORMATS]:
            from ultralytics.data.split import split_classify_dataset

            LOGGER.info(f"Found {len(image_files)} images in subdirectories. Attempting to split...")
            data_dir = split_classify_dataset(data_dir, train_ratio=0.8)
            train_set = data_dir / "train"
        else:
            raise FileNotFoundError(f"No images found in {data_dir} or its subdirectories.")
    val_set = (
        data_dir / "val"
        if (data_dir / "val").exists()
        else data_dir / "validation"
        if (data_dir / "validation").exists()
        else data_dir / "valid"
        if (data_dir / "valid").exists()
        else None
    )  # data/test or data/val
    test_set = data_dir / "test" if (data_dir / "test").exists() else None  # data/val or data/test
    if split == "val" and not val_set:
        LOGGER.warning("Dataset 'split=val' not found, using 'split=test' instead.")
        val_set = test_set
    elif split == "test" and not test_set:
        LOGGER.warning("Dataset 'split=test' not found, using 'split=val' instead.")
        test_set = val_set

    if (ndjson_names := data_dir / ".ndjson.yaml").is_file():
        names = YAML.load(ndjson_names)["names"]
    else:
        names = dict(enumerate(sorted(x.name for x in (data_dir / "train").iterdir() if x.is_dir())))
    nc = len(names)

    # Print to console
    for k, v in {"train": train_set, "val": val_set, "test": test_set}.items():
        prefix = f"{colorstr(f'{k}:')} {v}..."
        if v is None:
            LOGGER.info(prefix)
        else:
            files = [path for path in v.rglob("*.*") if path.suffix[1:].lower() in IMG_FORMATS]
            nf = len(files)  # number of files
            nd = len({file.parent for file in files})  # number of directories
            if nf == 0:
                if k == "train":
                    raise FileNotFoundError(f"{dataset} '{k}:' no training images found")
                else:
                    LOGGER.warning(f"{prefix} found {nf} images in {nd} classes (no images found)")
            elif nd != nc and not ndjson_names.is_file():
                LOGGER.error(f"{prefix} found {nf} images in {nd} classes (requires {nc} classes, not {nd})")
            else:
                class_count = f"{nd}/{nc}" if ndjson_names.is_file() else nd
                LOGGER.info(f"{prefix} found {nf} images in {class_count} classes ✅ ")

    return {"train": train_set, "val": val_set, "test": test_set, "nc": nc, "names": names, "channels": 3}


def compress_one_image(f: str | Path, f_new: str | Path | None = None, max_dim: int = 1920, quality: int = 50):
    """Compress a single image file to reduced size while preserving its aspect ratio.

    The image is saved as JPEG using the Python Imaging Library (PIL), falling back to OpenCV if PIL fails. If the input
    image is smaller than the maximum dimension, it will not be resized.

    Args:
        f (str | Path): The path to the input image file.
        f_new (str | Path, optional): The path to the output image file. If not specified, the input file will be
            overwritten.
        max_dim (int, optional): The maximum dimension (width or height) of the output image.
        quality (int, optional): The image compression quality as a percentage.

    Examples:
        >>> from pathlib import Path
        >>> from ultralytics.data.utils import compress_one_image
        >>> for f in Path("path/to/dataset").rglob("*.jpg"):
        ...     compress_one_image(f)
    """
    try:  # use PIL
        Image.MAX_IMAGE_PIXELS = None  # Fix DecompressionBombError, allow optimization of image > ~178.9 million pixels
        im = _exif_jpeg(f)  # JPEG save drops EXIF, so bake the orientation into the pixels
        if im.mode in {"RGBA", "LA"}:  # Convert to RGB if needed (for JPEG)
            im = im.convert("RGB")
        r = max_dim / max(im.height, im.width)  # ratio
        if r < 1.0:  # image too large
            im = im.resize((int(im.width * r), int(im.height * r)))
        _replace_image(f_new or f, lambda tmp: im.save(tmp, "JPEG", quality=quality, optimize=True))
    except Exception as e:  # use OpenCV
        LOGGER.warning(f"Image compression PIL failure {f}: {e}")
        im = cv2.imread(str(f))
        im_height, im_width = im.shape[:2]
        r = max_dim / max(im_height, im_width)  # ratio
        if r < 1.0:  # image too large
            im = cv2.resize(im, (int(im_width * r), int(im_height * r)), interpolation=cv2.INTER_AREA)
        _replace_image(f_new or f, lambda tmp: cv2.imwrite(tmp, im))


def load_dataset_cache_file(path: Path) -> dict:
    """Load an Ultralytics *.cache dictionary from path.

    Args:
        path (Path): Path to the *.cache file.

    Returns:
        (dict): The loaded cache dictionary.
    """
    import gc

    gc.disable()  # reduce pickle load time https://github.com/ultralytics/ultralytics/pull/1585
    try:
        return np.load(str(path), allow_pickle=True).item()  # load dict
    finally:
        gc.enable()  # also when loading raises, e.g. no cache file yet


def save_dataset_cache_file(prefix: str, path: Path, x: dict, version: str):
    """Save an Ultralytics dataset *.cache dictionary x to path.

    Args:
        prefix (str): Prefix for log messages.
        path (Path): Path to save the *.cache file.
        x (dict): Cache dictionary to save. A 'version' key is added in place.
        version (str): Cache version string.
    """
    x["version"] = version  # add cache version
    if is_dir_writeable(path.parent):
        if path.exists():
            path.unlink()  # remove *.cache file if exists
        try:
            with open(str(path), "wb") as file:  # context manager here fixes windows async np.save bug
                np.save(file, x)
            LOGGER.info(f"{prefix}New cache created: {path}")
        except Exception as e:
            Path(path).unlink(missing_ok=True)  # remove partially written file
            LOGGER.warning(f"{prefix}Failed to save cache to {path}: {e}")
    else:
        LOGGER.warning(f"{prefix}Cache directory {path.parent} is not writable, cache not saved.")


def add_polygon_background(data: dict) -> dict:
    """Set up the background class for polygon-based semantic datasets without 'masks_dir'.

    - nc > 1: appends a 'background' class at id=nc and bumps data['nc'] to nc+1; polygon cls values are kept as
    foreground ids.
    - nc == 1: keeps nc=1 (binary segmentation). Polygon rasterization yields a {0=bg, 1=fg} mask regardless of the
    label cls value.

    The data dictionary is modified in place and marked so repeated calls are no-ops.

    Args:
        data (dict): Dataset configuration dictionary.

    Returns:
        (dict): The updated dataset configuration dictionary, with 'bg_class_idx' set.
    """
    if data.get("masks_dir") or data.get("_polygon_bg_added"):
        return data
    nc = int(data.get("nc") or len(data.get("names") or {}))
    if nc == 1:  # binary: bg=0, fg=1 (implicit); model uses BCE on a single output channel
        data["bg_class_idx"] = 0
    else:
        names = dict(data.get("names") or {})
        names[nc] = "background"
        data["bg_class_idx"] = nc
        data["nc"] = nc + 1
        data["names"] = names
    data["_polygon_bg_added"] = True
    return data
