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

from __future__ import annotations

import json
from collections import defaultdict
from itertools import repeat
from multiprocessing.pool import ThreadPool
from pathlib import Path
from typing import Any

import cv2
import numpy as np
import torch
from PIL import Image
from torch.utils.data import ConcatDataset

from ultralytics.utils import LOCAL_RANK, LOGGER, NUM_THREADS, TQDM, IterableSimpleNamespace, colorstr
from ultralytics.utils.instance import Instances
from ultralytics.utils.ops import resample_segments, segments2boxes
from ultralytics.utils.patches import PIL_FALLBACK_SUFFIXES, imread, imread_unicode
from ultralytics.utils.torch_utils import TORCHVISION_0_18

from .augment import (
    Compose,
    DepthFormat,
    Format,
    LetterBox,
    RandomLoadText,
    SemanticFormat,
    classify_augmentations,
    classify_transforms,
    v8_transforms,
)
from .base import BaseDataset
from .converter import merge_multi_segment
from .utils import (
    HELP_URL,
    IMG_FORMATS,
    check_file_speeds,
    get_hash,
    get_split_fraction,
    img2label_paths,
    load_dataset_cache_file,
    load_depth,
    polygons2masks_overlap,
    save_dataset_cache_file,
    verify_image,
    verify_image_depth,
    verify_image_label,
    verify_image_mask,
)

# Ultralytics dataset *.cache version, >= 1.0.0 for Ultralytics YOLO models. Shared by every dataset type: a bump
# rescans all users' caches, so scope task-specific scan changes to that dataset's get_cache_hash() instead
DATASET_CACHE_VERSION = "1.0.10"  # EXIF-rotated image shapes now match the decoded image


class YOLODataset(BaseDataset):
    """Dataset class for loading object detection and/or segmentation labels in YOLO format.

    This class supports loading data for object detection, instance segmentation, pose estimation, and oriented bounding
    box (OBB) tasks using the YOLO format.

    Attributes:
        format_class (type[Format]): Formatter appended by build_transforms; subclasses override it per task.
        use_segments (bool): Indicates if segmentation masks should be used.
        use_keypoints (bool): Indicates if keypoints should be used for pose estimation.
        use_obb (bool): Indicates if oriented bounding boxes should be used.
        data (dict): Dataset configuration dictionary.

    Methods:
        cache_labels: Cache dataset labels, check images and read shapes.
        get_labels: Return list of label dictionaries for YOLO training.
        get_label_files: Return companion label files for the dataset's images.
        verify_args: Return the per-image verification function and its arguments.
        result_to_label: Convert one verification result into a label dict.
        verify_labels: Check box/segment consistency of the loaded labels.
        get_cache_hash: Return the hash used to validate a label cache.
        scan_summary: Return a one-line summary of scan counters.
        build_transforms: Build and append transforms to the list.
        build_text_transforms: Insert text augmentation for text-based subclasses.
        close_mosaic: Disable mosaic, copy_paste, mixup and cutmix augmentations and build transformations.
        update_labels_info: Update label format for different tasks.
        collate_fn: Collate data samples into batches.

    Examples:
        >>> dataset = YOLODataset(img_path="path/to/images", data={"names": {0: "person"}}, task="detect")
        >>> dataset.get_labels()
    """

    format_class = Format

    def __init__(self, *args, data: dict, task: str = "detect", **kwargs):
        """Initialize the YOLODataset.

        Args:
            data (dict): Dataset configuration dictionary.
            task (str): Task type, one of 'detect', 'segment', 'pose', or 'obb'.
            *args (Any): Additional positional arguments for the parent class.
            **kwargs (Any): Additional keyword arguments for the parent class.

        Raises:
            ValueError: If task is 'pose' and data['kpt_shape'] is missing or invalid.
        """
        self.use_segments = task == "segment"
        self.use_keypoints = task == "pose"
        self.use_obb = task == "obb"
        self.data = data
        nkpt, ndim = self.data.get("kpt_shape", (0, 0))
        if self.use_keypoints and (nkpt <= 0 or ndim not in {2, 3}):  # checked before the label cache is consulted
            raise ValueError(
                "'kpt_shape' in data.yaml missing or incorrect. Should be a list with [number of "
                "keypoints, number of dims (2 for x,y or 3 for x,y,visible)], i.e. 'kpt_shape: [17, 3]'"
            )
        super().__init__(*args, channels=self.data.get("channels", 3), **kwargs)

    def cache_labels(self, path: Path = Path("./labels.cache")) -> dict:
        """Cache dataset labels, check images and read shapes.

        This is the shared scanning skeleton for file-based datasets; subclasses customize it through the
        `get_label_files`, `get_cache_hash`, `verify_args`, `result_to_label` and `scan_summary` hooks instead of
        duplicating this method.

        Args:
            path (Path): Path where to save the cache file.

        Returns:
            (dict): Dictionary containing cached labels and related information.
        """
        x = {"labels": []}
        nm, nf, ne, nc, msgs = 0, 0, 0, 0, []  # number missing, found, empty, corrupt, messages
        desc = f"{self.prefix}Scanning {path.parent / path.stem}..."
        total = len(self.im_files)
        with ThreadPool(NUM_THREADS) as pool:
            func, iterable = self.verify_args()
            results = pool.imap(func=func, iterable=iterable)
            pbar = TQDM(results, desc=desc, total=total)
            for result in pbar:
                label, nm_f, nf_f, ne_f, nc_f, msg = self.result_to_label(result)
                nm += nm_f
                nf += nf_f
                ne += ne_f
                nc += nc_f
                if label is not None:
                    x["labels"].append(label)
                if msg:
                    msgs.append(msg)
                pbar.desc = f"{desc} {self.scan_summary(nf, nm, ne, nc)}"
            pbar.close()

        if msgs:
            LOGGER.info("\n".join(msgs))
        if nf == 0:
            if self.augment:  # training requires labels; unlabeled val splits (e.g. COCO test-dev) only warn
                raise ValueError(f"{self.prefix}No labels found in {path}. {HELP_URL}")
            LOGGER.warning(f"{self.prefix}No labels found in {path}. {HELP_URL}")
        x["hash"] = self.get_cache_hash()
        x["results"] = nf, nm, ne, nc, total
        x["msgs"] = msgs  # warnings
        if x["labels"]:
            save_dataset_cache_file(self.prefix, path, x, DATASET_CACHE_VERSION)
        return x

    def get_label_files(self) -> list[str]:
        """Return the companion label files for the dataset's images, storing them on the instance.

        Returns:
            (list[str]): List of label file paths.
        """
        self.label_files = img2label_paths(self.im_files)
        return self.label_files

    def get_cache_hash(self) -> str:
        """Return the hash used to validate a label cache against the current dataset files and scan settings.

        Returns:
            (str): Dataset cache hash.
        """
        # add_polygon_background() class is not a label class, so segment and semantic share one cache
        nc = self.data.get("bg_class_idx") or len(self.data["names"])
        scan_args = (self.use_keypoints, nc, self.data.get("kpt_shape"), self.single_cls)
        return get_hash(self.label_files + self.im_files + [str(scan_args)])

    def scan_summary(self, nf: int, nm: int, ne: int, nc: int) -> str:
        """Return a one-line summary of scan counters for progress bars and cache logs.

        Args:
            nf (int): Number of found images.
            nm (int): Number of missing labels.
            ne (int): Number of empty labels.
            nc (int): Number of corrupt images.

        Returns:
            (str): Scan summary message.
        """
        return f"{nf} images, {nm + ne} backgrounds, {nc} corrupt"

    def verify_args(self) -> tuple:
        """Return the per-image verification function and its argument iterable used by `cache_labels`.

        Returns:
            (tuple): (verify function, zipped argument iterable) for ThreadPool.imap.
        """
        nkpt, ndim = self.data.get("kpt_shape", (0, 0))
        return verify_image_label, zip(
            self.im_files,
            self.label_files,
            repeat(self.prefix),
            repeat(self.use_keypoints),
            repeat(self.data.get("bg_class_idx") or len(self.data["names"])),  # label classes, no semantic background
            repeat(nkpt),
            repeat(ndim),
            repeat(self.single_cls),
        )

    def result_to_label(self, result: list) -> tuple[dict | None, int, int, int, int, str]:
        """Convert one verification result into a label dict and scan counter increments.

        Args:
            result (list): One result from the verification function returned by `verify_args`.

        Returns:
            (tuple): (label dict or None, missing, found, empty, corrupt, message).
        """
        im_file, lb, shape, segments, keypoint, nm_f, nf_f, ne_f, nc_f, msg = result
        label = (
            {
                "im_file": im_file,
                "shape": shape,
                "cls": lb[:, 0:1],  # n, 1
                "bboxes": lb[:, 1:],  # n, 4
                "segments": segments,
                "keypoints": keypoint,
                "normalized": True,
                "bbox_format": "xywh",
            }
            if im_file
            else None
        )
        return label, nm_f, nf_f, ne_f, nc_f, msg

    def verify_labels(self, labels: list[dict], cache_path: Path) -> None:
        """Check that the dataset is all boxes or all segments, removing mixed segments if necessary.

        Args:
            labels (list[dict]): List of label dictionaries.
            cache_path (Path): Path of the dataset cache file, used in warning messages.
        """
        # Check if the dataset is all boxes or all segments
        lengths = ((len(lb["cls"]), len(lb["bboxes"]), len(lb["segments"])) for lb in labels)
        len_cls, len_boxes, len_segments = (sum(x) for x in zip(*lengths))
        if (self.use_segments or self.use_obb) and len_boxes != len_segments:
            task = "OBB" if self.use_obb else "Segment"
            raise ValueError(
                f"{task} dataset requires equal numbers of boxes and segments, but got len(segments) = "
                f"{len_segments}, len(boxes) = {len_boxes}. Please supply {'an OBB' if self.use_obb else 'a segment'} "
                "dataset, not a detect dataset."
            )
        if len_segments and len_boxes != len_segments:
            LOGGER.warning(
                f"Box and segment counts should be equal, but got len(segments) = {len_segments}, "
                f"len(boxes) = {len_boxes}. To resolve this only boxes will be used and all segments will be removed. "
                "To avoid this please supply either a detect or segment dataset, not a detect-segment mixed dataset."
            )
            for lb in labels:
                lb["segments"] = []
        if len_cls == 0:
            LOGGER.warning(f"Labels are missing or empty in {cache_path}, training may not work correctly. {HELP_URL}")

    def _load_or_scan_cache(self, cache_path: Path, cache_hash: str) -> tuple[dict, bool]:
        """Load a dataset cache file if it matches the current version and hash, otherwise rescan and rebuild it.

        Args:
            cache_path (Path): Path of the cache file.
            cache_hash (str): Expected hash of the dataset files.

        Returns:
            (tuple): (cache dict, True if a valid existing cache file was loaded).
        """
        try:
            cache, exists = load_dataset_cache_file(cache_path), True  # attempt to load a *.cache file
            assert cache["version"] == DATASET_CACHE_VERSION  # matches current version
            assert cache["hash"] == cache_hash  # identical hash
        except Exception:  # missing, stale, or unreadable (e.g. truncated) cache
            cache, exists = self.cache_labels(cache_path), False  # run cache ops
        return cache, exists

    def get_labels(self) -> list[dict]:
        """Return list of label dictionaries for YOLO training.

        This method loads labels from disk or cache, verifies their integrity, and prepares them for training.

        Returns:
            (list[dict]): List of label dictionaries, each containing information about an image and its annotations.

        Raises:
            RuntimeError: If no valid images are found.
        """
        label_files = self.get_label_files()
        cache_path = Path(label_files[0]).parent.with_suffix(".cache")
        cache, exists = self._load_or_scan_cache(cache_path, self.get_cache_hash())

        # Display cache
        nf, nm, ne, nc, n = cache.pop("results")  # found, missing, empty, corrupt, total
        if exists and LOCAL_RANK in {-1, 0}:
            d = f"Scanning {cache_path}... {self.scan_summary(nf, nm, ne, nc)}"
            TQDM(None, desc=self.prefix + d, total=n, initial=n)  # display results
            if cache["msgs"]:
                LOGGER.info("\n".join(cache["msgs"]))  # display warnings

        # Read cache
        labels = cache["labels"]
        if not labels:
            issues = "\n  ".join(sorted(set(cache["msgs"]))) or "no error details"
            raise RuntimeError(f"No valid images found in {cache_path}.\n  {issues}\n{HELP_URL}")
        [cache.pop(k) for k in ("hash", "version", "msgs")]  # remove items
        self.im_files = [lb["im_file"] for lb in labels]  # update im_files
        self.verify_labels(labels, cache_path)
        return labels

    def build_transforms(self, hyp: IterableSimpleNamespace) -> Compose:
        """Build and append transforms to the list.

        Args:
            hyp (IterableSimpleNamespace): Hyperparameters for transforms.

        Returns:
            (Compose): Composed transforms.
        """
        if self.augment:
            hyp.mosaic = hyp.mosaic if self.augment and not self.rect else 0.0
            hyp.mixup = hyp.mixup if self.augment and not self.rect else 0.0
            hyp.cutmix = hyp.cutmix if self.augment and not self.rect else 0.0
            transforms = v8_transforms(self, self.imgsz, hyp)
            if self.format_class is SemanticFormat:  # masks rasterize from self.labels; only these read polygons
                self.use_segments = bool(hyp.copy_paste or hyp.cutmix or getattr(hyp, "augmentations", None))
        else:
            transforms = Compose([LetterBox(new_shape=(self.imgsz, self.imgsz), scaleup=False)])
        transforms.append(
            self.format_class(
                bbox_format="xywh",
                normalize=True,
                return_mask=self.use_segments,
                return_keypoint=self.use_keypoints,
                return_obb=self.use_obb,
                batch_idx=True,
                mask_ratio=hyp.mask_ratio,
                mask_overlap=hyp.overlap_mask,
                bgr=hyp.bgr if self.augment else 0.0,  # only affect training.
            )
        )
        return transforms

    def build_text_transforms(self, transforms: Compose, max_samples: int) -> Compose:
        """Insert text augmentation for text-based subclasses providing `category_freq`.

        Args:
            transforms (Compose): Transforms composed by build_transforms.
            max_samples (int): Maximum number of text samples per image.

        Returns:
            (Compose): Transforms with RandomLoadText inserted before Format when augmenting.
        """
        if self.augment:
            # NOTE: hard-coded the args for now.
            # NOTE: this implementation is different from official yoloe,
            # the strategy of selecting negative is restricted in one dataset,
            # while official pre-saved neg embeddings from all datasets at once.
            transform = RandomLoadText(
                max_samples=min(max_samples, 80),
                padding=True,
                padding_value=self._get_neg_texts(self.category_freq),
            )
            transforms.insert(-1, transform)
        return transforms

    @staticmethod
    def _get_neg_texts(category_freq: dict) -> list[str]:
        """Get negative text samples with frequency above the dataset threshold."""
        threshold = min(max(category_freq.values()), 100)
        return [k for k, v in category_freq.items() if v >= threshold]

    def close_mosaic(self, hyp: IterableSimpleNamespace) -> None:
        """Disable mosaic, copy_paste, mixup and cutmix augmentations by setting their values to 0.0.

        Args:
            hyp (IterableSimpleNamespace): Hyperparameters for transforms.
        """
        hyp.mosaic = 0.0
        hyp.copy_paste = 0.0
        hyp.mixup = 0.0
        hyp.cutmix = 0.0
        self.transforms = self.build_transforms(hyp)

    def update_labels_info(self, label: dict) -> dict:
        """Update label format for different tasks.

        Args:
            label (dict): Label dictionary containing bboxes, segments, keypoints, etc.

        Returns:
            (dict): Updated label dictionary with instances.

        Notes:
            cls is not with bboxes now, classification and semantic segmentation need an independent cls label
            Can also support classification and semantic segmentation by adding or removing dict keys there.
        """
        bboxes = label.pop("bboxes")
        segments = label.pop("segments", [])
        keypoints = label.pop("keypoints", None)
        bbox_format = label.pop("bbox_format")
        normalized = label.pop("normalized")

        # NOTE: do NOT resample oriented boxes
        segment_resamples = 100 if self.use_obb else 1000
        if len(segments) > 0 and (self.use_segments or self.format_class is not SemanticFormat):
            # make sure segments interpolate correctly if original length is greater than segment_resamples
            max_len = max(len(s) for s in segments)
            segment_resamples = (max_len + 1) if segment_resamples < max_len else segment_resamples
            # list[np.array(segment_resamples, 2)] * num_samples
            segments = np.stack(resample_segments(segments, n=segment_resamples), axis=0)
        else:
            segments = np.zeros((0, segment_resamples, 2), dtype=np.float32)
        label["instances"] = Instances(bboxes, segments, keypoints, bbox_format=bbox_format, normalized=normalized)
        return label

    @staticmethod
    def collate_fn(batch: list[dict]) -> dict:
        """Collate data samples into batches.

        Args:
            batch (list[dict]): List of dictionaries containing sample data.

        Returns:
            (dict): Collated batch with stacked tensors.
        """
        new_batch = {}
        batch = [dict(sorted(b.items())) for b in batch]  # make sure the keys are in the same order
        keys = batch[0].keys()
        values = list(zip(*[list(b.values()) for b in batch]))
        for i, k in enumerate(keys):
            value = values[i]
            if k in {"img", "text_feats", "semantic_mask", "sem_masks", "depth"}:
                value = torch.stack(value, 0)
            elif k == "visuals":
                value = torch.nn.utils.rnn.pad_sequence(value, batch_first=True)
            if k in {"masks", "keypoints", "bboxes", "cls", "segments", "obb"}:
                value = torch.cat(value, 0)
            new_batch[k] = value
        if "batch_idx" in new_batch:
            new_batch["batch_idx"] = list(new_batch["batch_idx"])
            for i in range(len(new_batch["batch_idx"])):
                new_batch["batch_idx"][i] += i  # add target image index for build_targets()
            new_batch["batch_idx"] = torch.cat(new_batch["batch_idx"], 0)
        return new_batch


class DepthDataset(YOLODataset):
    """Dataset for monocular depth estimation with paired RGB + depth map loading.

    Extends YOLODataset to load depth ground truth maps alongside RGB images. Depth maps are stored as PNG or NPY files
    in a parallel directory structure (images/train/*.jpg → depth/train/*.{png,npy}).

    Examples:
        >>> dataset = DepthDataset(img_path="/data/nyu/images/train", data={"nc": 1})
    """

    format_class = DepthFormat

    def _depth_path_for(self, im_file: str) -> str:
        """Map an image path to its companion PNG or NPY depth target."""
        parts = list(Path(im_file).parts)
        for i in range(len(parts) - 1, -1, -1):
            if parts[i] == "images":
                parts[i] = "depth"
                break
        path = Path(*parts).with_suffix(".png")
        return str(path if path.is_file() else path.with_suffix(".npy"))

    def get_label_files(self) -> list[str]:
        """Return the depth paths paired with the dataset's images.

        Returns:
            (list[str]): List of depth file paths.
        """
        self.depth_files_by_image = {f: self._depth_path_for(f) for f in self.im_files}
        self.depth_files = list(self.depth_files_by_image.values())
        return self.depth_files

    def get_cache_hash(self) -> str:
        """Return a hash over the paired depth and image files.

        Returns:
            (str): Dataset cache hash.
        """
        return get_hash(self.depth_files + self.im_files + [str(self.data.get("depth_scale", 1000))])

    def scan_summary(self, nf: int, nm: int, ne: int, nc: int) -> str:
        """Return a one-line summary of image-depth scan counters."""
        return f"{nf} images, {nm} missing depth, {nc} corrupt"

    def verify_args(self) -> tuple:
        """Return the depth verification function and its argument iterable."""
        return verify_image_depth, zip(
            self.im_files, self.depth_files, repeat(self.prefix), repeat(self.data.get("depth_scale", 1000))
        )

    def result_to_label(self, result: tuple) -> tuple[dict | None, int, int, int, int, str]:
        """Convert one verify_image_depth result into a label dict and scan counter increments."""
        im_file, shape, nf_f, nm_f, nc_f, msg = result
        label = (
            {
                "im_file": im_file,
                "shape": shape,
                "cls": np.array([], dtype=np.float32),
                "bboxes": np.zeros((0, 4), dtype=np.float32),
                "segments": [],
                "normalized": True,
                "bbox_format": "xywh",
            }
            if im_file
            else None
        )
        return label, nm_f, nf_f, 0, nc_f, msg

    def verify_labels(self, labels: list[dict], cache_path: Path) -> None:
        """Skip box and segment checks; depth datasets carry no box or segment annotations."""

    def _load_depth(self, index: int) -> np.ndarray:
        """Return the native-resolution depth map for an image."""
        return load_depth(self.depth_files_by_image[self.im_files[index]], self.data.get("depth_scale", 1000))

    def get_image_and_label(self, index: int) -> dict[str, Any]:
        """Load image, label, and depth map for the given index."""
        label = super().get_image_and_label(index)
        h, w = label["resized_shape"]
        depth = self._load_depth(index)
        if depth.shape[:2] != (h, w):
            depth = cv2.resize(depth, (w, h), interpolation=cv2.INTER_NEAREST)
        label["depth"] = depth
        return label

    def build_transforms(self, hyp: IterableSimpleNamespace) -> Compose:
        """Build transforms for depth estimation.

        Args:
            hyp (IterableSimpleNamespace): Hyperparameters.

        Returns:
            (Compose): Composed transforms.
        """
        # NOTE: For now following arguments are not supported
        hyp.mosaic = hyp.mixup = hyp.cutmix = hyp.copy_paste = 0.0
        transforms = super().build_transforms(hyp)
        if not self.augment:
            # stretch the image instead of padding
            transforms[-2] = LetterBox(new_shape=(self.imgsz, self.imgsz), scale_fill=True)
        return transforms


class YOLOMultiModalDataset(YOLODataset):
    """Dataset class for loading object detection and/or segmentation labels in YOLO format with multi-modal support.

    This class extends YOLODataset to add text information for multi-modal model training, enabling models to process
    both image and text data.

    Methods:
        update_labels_info: Add text information for multi-modal model training.
        build_transforms: Enhance data transformations with text augmentation.

    Examples:
        >>> dataset = YOLOMultiModalDataset(img_path="path/to/images", data={"names": {0: "person"}}, task="detect")
        >>> sample = dataset[0]
        >>> print(sample.keys())  # Should include 'texts'
    """

    def update_labels_info(self, label: dict) -> dict:
        """Add text information for multi-modal model training.

        Args:
            label (dict): Label dictionary containing bboxes, segments, keypoints, etc.

        Returns:
            (dict): Updated label dictionary with instances and texts.
        """
        labels = super().update_labels_info(label)
        # NOTE: some categories are concatenated with its synonyms by `/`.
        # NOTE: and `RandomLoadText` would randomly select one of them if there are multiple words.
        labels["texts"] = [v.split("/") for _, v in self.data["names"].items()]

        return labels

    def build_transforms(self, hyp: IterableSimpleNamespace) -> Compose:
        """Enhance data transformations with text augmentation for multi-modal training.

        Args:
            hyp (IterableSimpleNamespace): Hyperparameters for transforms.

        Returns:
            (Compose): Composed transforms including text augmentation if applicable.
        """
        return self.build_text_transforms(super().build_transforms(hyp), self.data["nc"])

    @property
    def category_names(self):
        """Return category names for the dataset.

        Returns:
            (set[str]): Set of class names.
        """
        names = self.data["names"].values()
        return {n.strip() for name in names for n in name.split("/")}  # category names

    @property
    def category_freq(self):
        """Return frequency of each category in the dataset."""
        texts = [v.split("/") for v in self.data["names"].values()]
        category_freq = defaultdict(int)
        for label in self.labels:
            for c in label["cls"].squeeze(-1):  # to check
                text = texts[int(c)]
                for t in text:
                    t = t.strip()
                    category_freq[t] += 1
        # a background-only dataset sees no class, leaving every class an equally valid negative
        return category_freq or dict.fromkeys((t.strip() for text in texts for t in text), 0)


class GroundingDataset(YOLODataset):
    """Dataset class for object detection tasks using annotations from a JSON file in grounding format.

    This dataset is designed for grounding tasks where annotations are provided in a JSON file rather than the standard
    YOLO format text files.

    Attributes:
        json_file (str): Path to the JSON file containing annotations.

    Methods:
        get_labels: Load annotations from a JSON file and prepare them for training.
        build_transforms: Configure augmentations for training with optional text loading.

    Examples:
        >>> dataset = GroundingDataset(img_path="path/to/images", json_file="annotations.json", task="detect")
        >>> len(dataset)  # Number of valid images with annotations
    """

    def __init__(self, *args, json_file: str, task: str = "detect", max_samples: int = 80, **kwargs):
        """Initialize a GroundingDataset for object detection.

        Args:
            json_file (str): Path to the JSON file containing annotations.
            task (str): Must be 'detect' or 'segment' for GroundingDataset.
            max_samples (int): Maximum number of text samples per image for text augmentation.
            *args (Any): Additional positional arguments for the parent class.
            **kwargs (Any): Additional keyword arguments for the parent class.
        """
        assert task in {"detect", "segment"}, "GroundingDataset currently only supports `detect` and `segment` tasks"
        self.json_file = json_file
        self.max_samples = max_samples
        super().__init__(*args, task=task, data={"channels": 3}, **kwargs)

    def get_img_files(self, img_path: str) -> list[str]:
        """Return every image under `img_path`; the annotations, not `fraction`, decide which ones are used."""
        self.fraction = 1.0  # a truncated inventory would leave later images outside the cache key
        self.scan_files = super().get_img_files(img_path)
        return self.scan_files

    def get_cache_hash(self) -> str:
        """Return a hash over the annotation file and images scanned against it."""
        return get_hash([self.json_file, *self.scan_files])

    def _verify_instance_counts(self, labels: list[dict[str, Any]]) -> None:
        """Verify instance counts for known grounding datasets."""
        expected_counts = {
            "final_mixed_train_no_coco_segm": 3662412,
            "final_mixed_train_no_coco": 3681235,
            "final_flickr_separateGT_train_segm": 638214,
            "final_flickr_separateGT_train": 640704,
        }

        instance_count = sum(label["bboxes"].shape[0] for label in labels)
        for data_name, count in expected_counts.items():
            if data_name in self.json_file:
                assert instance_count == count, f"'{self.json_file}' has {instance_count} instances, expected {count}."
                return
        LOGGER.warning(f"Skipping instance count verification for unrecognized dataset '{self.json_file}'")

    def cache_labels(self, path: Path = Path("./labels.cache")) -> dict[str, Any]:
        """Load annotations from a JSON file, filter, and normalize bounding boxes for each image.

        Args:
            path (Path): Path where to save the cache file.

        Returns:
            (dict[str, Any]): Dictionary containing cached labels and related information.
        """
        x = {"labels": []}
        LOGGER.info("Loading annotation file...")
        with open(self.json_file) as f:
            annotations = json.load(f)
        images = {f"{x['id']:d}": x for x in annotations["images"]}
        img_to_anns = defaultdict(list)
        for ann in annotations["annotations"]:
            img_to_anns[ann["image_id"]].append(ann)
        dropped = False
        for img_id, anns in TQDM(img_to_anns.items(), desc=f"Reading annotations {self.json_file}"):
            img = images[f"{img_id:d}"]
            h, w, f = img["height"], img["width"], img["file_name"]
            im_file = Path(self.img_path) / f
            if not im_file.exists():
                continue
            bboxes = []
            segments = []
            segmented = False
            cat2id = {}
            texts = []
            for ann in anns:
                if ann["iscrowd"]:
                    continue
                box = np.array(ann["bbox"], dtype=np.float32)
                box[:2] += box[2:] / 2
                box[[0, 2]] /= float(w)
                box[[1, 3]] /= float(h)
                if box[2] <= 0 or box[3] <= 0:
                    continue

                caption = img["caption"]
                cat_name = " ".join([caption[t[0] : t[1]] for t in ann["tokens_positive"]]).lower().strip()
                if not cat_name:
                    continue

                if cat_name not in cat2id:
                    cat2id[cat_name] = len(cat2id)
                    texts.append([cat_name])
                cls = cat2id[cat_name]  # class
                box = [cls, *box.tolist()]
                if box not in bboxes:
                    bboxes.append(box)
                    raw_seg = ann.get("segmentation")
                    segmented |= raw_seg is not None
                    seg = raw_seg if isinstance(raw_seg, list) else []
                    polygons = [
                        p
                        for p in seg
                        if isinstance(p, list)
                        and len(p) >= 6
                        and not len(p) % 2
                        and all(isinstance(c, (int, float)) for c in p)
                    ]
                    dropped |= bool(raw_seg) and (not isinstance(raw_seg, list) or len(polygons) < len(seg))
                    if not polygons:  # keep one segment per box so an image mixing the two kinds stays aligned
                        cx, cy, bw, bh = box[1:]
                        x1, y1, x2, y2 = cx - bw / 2, cy - bh / 2, cx + bw / 2, cy + bh / 2
                        segments.append([cls, x1, y1, x2, y1, x2, y2, x1, y2])  # segments2boxes returns the box
                        continue
                    elif len(polygons) > 1:
                        s = merge_multi_segment(polygons)
                        s = (np.concatenate(s, axis=0) / np.array([w, h], dtype=np.float32)).reshape(-1).tolist()
                    else:
                        s = [j for i in polygons for j in i]  # all segments concatenated
                        s = (
                            (np.array(s, dtype=np.float32).reshape(-1, 2) / np.array([w, h], dtype=np.float32))
                            .reshape(-1)
                            .tolist()
                        )
                    segments.append([cls, *s])
            lb = np.array(bboxes, dtype=np.float32) if len(bboxes) else np.zeros((0, 5), dtype=np.float32)

            if segmented:
                segments = [np.array(x[1:], dtype=np.float32).reshape(-1, 2) for x in segments]  # (cls, xy1...)
                lb[:, 1:] = segments2boxes(segments)  # boxes follow the polygons
            else:
                segments = []  # no annotation carried a segmentation, so store no masks

            x["labels"].append(
                {
                    "im_file": im_file,
                    "shape": (h, w),
                    "cls": lb[:, 0:1],  # n, 1
                    "bboxes": lb[:, 1:],  # n, 4
                    "segments": segments,
                    "normalized": True,
                    "bbox_format": "xywh",
                    "texts": texts,
                }
            )
        if dropped:
            LOGGER.warning(
                f"{self.json_file}: ignored segmentations that are not polygon point lists, such as RLE masks. "
                "Annotations left without a polygon use a segment shaped like their bounding box."
            )
        x["hash"] = self.get_cache_hash()
        save_dataset_cache_file(self.prefix, path, x, DATASET_CACHE_VERSION)
        return x

    def get_labels(self) -> list[dict]:
        """Load labels from cache or generate them from JSON file.

        Returns:
            (list[dict]): List of label dictionaries, each containing information about an image and its annotations.
        """
        cache_path = Path(self.json_file).with_suffix(".cache")
        cache, _ = self._load_or_scan_cache(cache_path, self.get_cache_hash())
        [cache.pop(k) for k in ("hash", "version")]  # remove items
        labels = cache["labels"]
        if not labels:
            raise RuntimeError(f"No images from {self.json_file} found in {self.img_path}. {HELP_URL}")
        if not any(label["texts"] for label in labels):  # category_freq is empty, so negative texts cannot be built
            raise RuntimeError(
                f"No annotations in {self.json_file} survived filtering. Every one is iscrowd, resolves to an empty "
                f"caption span or has a zero-size box. {HELP_URL}"
            )
        self._verify_instance_counts(labels)
        self.im_files = [str(label["im_file"]) for label in labels]
        if LOCAL_RANK in {-1, 0}:
            LOGGER.info(f"Load {self.json_file} from cache file {cache_path}")
        return labels

    def build_transforms(self, hyp: IterableSimpleNamespace) -> Compose:
        """Configure augmentations for training with optional text loading.

        Args:
            hyp (IterableSimpleNamespace): Hyperparameters for transforms.

        Returns:
            (Compose): Composed transforms including text augmentation if applicable.
        """
        return self.build_text_transforms(super().build_transforms(hyp), self.max_samples)

    @property
    def category_names(self):
        """Return unique category names from the dataset."""
        return {t.strip() for label in self.labels for text in label["texts"] for t in text}

    @property
    def category_freq(self):
        """Return frequency of each category in the dataset."""
        category_freq = defaultdict(int)
        for label in self.labels:
            for text in label["texts"]:
                for t in text:
                    t = t.strip()
                    category_freq[t] += 1
        return category_freq


class YOLOConcatDataset(ConcatDataset):
    """Dataset as a concatenation of multiple datasets.

    This class is useful to assemble different existing datasets for YOLO training, ensuring they use the same collation
    function.

    Methods:
        collate_fn: Static method that collates data samples into batches using YOLODataset's collation function.

    Examples:
        >>> dataset1 = YOLODataset(...)
        >>> dataset2 = YOLODataset(...)
        >>> combined_dataset = YOLOConcatDataset([dataset1, dataset2])
    """

    @staticmethod
    def collate_fn(batch: list[dict]) -> dict:
        """Collate data samples into batches.

        Args:
            batch (list[dict]): List of dictionaries containing sample data.

        Returns:
            (dict): Collated batch with stacked tensors.
        """
        return YOLODataset.collate_fn(batch)

    def close_mosaic(self, hyp: IterableSimpleNamespace) -> None:
        """Disable mosaic, copy_paste, mixup and cutmix augmentations by setting their values to 0.0.

        Args:
            hyp (IterableSimpleNamespace): Hyperparameters for transforms.
        """
        for dataset in self.datasets:
            if not hasattr(dataset, "close_mosaic"):
                continue
            dataset.close_mosaic(hyp)


class SemanticDataset(YOLODataset):
    """Dataset for semantic segmentation with PNG mask labels.

    Expects a directory structure where each image has a corresponding PNG mask file with the same stem. Pixel values in
    masks represent class IDs, with 255 as the ignore label.

    The mask directory is specified in the dataset YAML via 'masks_dir' key, and mirrors the images/ directory structure
    (e.g., images/train/ -> masks/train/).

    Attributes:
        data (dict): Dataset configuration from YAML.
        mask_files (list[str]): List of mask file paths corresponding to images.
        include_class (np.ndarray | None): Class ids to keep per pixel (None keeps all).
        masks (dict[int, np.ndarray]): Resized masks of the images in the mosaic buffer, evicted with them.
        label_mapping (dict[int, int]): Mapping from raw mask ids to training ids from the dataset YAML 'label_mapping'
            key, where 255 is the ignore label.
        label_lut (np.ndarray): 256-entry lookup table applying label_mapping.
        inverse_lut (np.ndarray): 256-entry lookup table reverting label_mapping.
    """

    format_class = SemanticFormat

    def __init__(self, *args, data: dict, **kwargs):
        """Initialize SemanticDataset.

        Args:
            *args (Any): Additional positional arguments for the parent class.
            data (dict): Dataset configuration dictionary.
            **kwargs (Any): Additional keyword arguments for the parent class.
        """
        self.data = data
        self.label_mapping = self._parse_label_mapping(self.data.get("label_mapping"))
        self.label_lut, self.inverse_lut = self._build_label_luts()
        self.mask_files = []
        self.include_class = None
        self.masks = {}  # masks of the buffered images, evicted with the image buffer
        super().__init__(*args, data=data, **kwargs)

    def update_labels(self, include_class: list[int] | None) -> None:
        """Store the classes to keep per pixel; pixels of other classes are set to the ignore label (255) on load.

        Args:
            include_class (list[int], optional): List of classes to include. If None, all classes are included.

        Raises:
            NotImplementedError: If single_cls is True.
        """
        if self.single_cls:
            raise NotImplementedError(
                "'single_cls=True' is not supported for semantic segmentation: it forces a single-channel "
                "model but cannot collapse multi-class masks. Use a dataset with 'nc: 1' for binary "
                "(foreground/background) segmentation instead."
            )
        self.include_class = None if include_class is None else np.asarray(include_class, dtype=np.int32).reshape(-1)
        if self.include_class is not None and int(self.data.get("nc", 0)) == 1:
            LOGGER.warning(
                "'classes' filtering is ignored for single-class (binary) semantic segmentation: keeping only "
                "the sole class would discard all background supervision."
            )
            self.include_class = None

    def _parse_label_mapping(self, mapping):
        """Normalize label_mapping entries from dataset YAML into integer-to-integer ids."""
        if mapping is None:
            return {}
        if not isinstance(mapping, dict):
            raise TypeError(f"Expected 'label_mapping' to be a dict in dataset YAML, but got {type(mapping).__name__}.")

        normalized = {}
        for src, dst in mapping.items():
            src = int(src)
            if isinstance(dst, str):
                dst = dst.strip()
                dst = 255 if dst == "ignore_label" else int(dst)
            elif dst is None:
                dst = 255
            else:
                dst = int(dst)
            normalized[src] = dst
        return normalized

    def _build_label_luts(self) -> tuple[np.ndarray, np.ndarray]:
        """Build the 256-entry forward and inverse lookup tables for the dataset label mapping."""
        forward, inverse = np.arange(256, dtype=np.uint8), np.arange(256, dtype=np.uint8)
        for k, v in self.label_mapping.items():  # ids outside 0-255 never match a uint8 mask pixel
            if 0 <= k < 256:
                forward[k] = v
            if 0 <= v < 256:
                inverse[v] = k & 0xFF  # cityscapes maps -1; the inverse caller casts the result to uint8
        return forward, inverse

    def get_label_files(self) -> list[str]:
        """Return the mask PNG paths paired with the dataset's images.

        Returns:
            (list[str]): List of mask file paths.
        """
        self.mask_files = img2label_paths(self.im_files, label_dir=self.data.get("masks_dir", "masks"), suffix=".png")
        return self.mask_files

    def get_cache_hash(self) -> str:
        """Return a hash for semantic cache validation that also includes label_mapping changes.

        Returns:
            (str): Dataset cache hash.
        """
        mapping = json.dumps(self.label_mapping, sort_keys=True, separators=(",", ":"))
        return get_hash(self.im_files + self.mask_files + [f"label_mapping:{mapping}"])

    def scan_summary(self, nf: int, nm: int, ne: int, nc: int) -> str:
        """Return a one-line summary of image-mask scan counters."""
        return f"{nf} images, {nm} missing masks, {nc} corrupt"

    def verify_args(self) -> tuple:
        """Return the mask verification function and its argument iterable."""
        nc = len(self.data["names"])
        invalid = ((self.label_lut > max(nc - 1, 1)) & (self.label_lut != 255)).astype(np.uint8)  # nc=1 keeps {0, 1}
        return verify_image_mask, zip(self.im_files, self.mask_files, repeat(self.prefix), repeat(invalid))

    def result_to_label(self, result: tuple) -> tuple[dict | None, int, int, int, int, str]:
        """Convert one verify_image_mask result into a label dict and scan counter increments."""
        im_file, mask_file, shape, mode, nm_f, nf_f, nc_f, msg = result
        label = (
            {
                "im_file": im_file,
                "mask_file": mask_file,
                "shape": shape,
                "mode": mode,
                "cls": np.array([], dtype=np.float32),
                "bboxes": np.zeros((0, 4), dtype=np.float32),
                "segments": [],
                "normalized": True,
                "bbox_format": "xywh",
            }
            if im_file
            else None
        )
        return label, nm_f, nf_f, 0, nc_f, msg

    def verify_labels(self, labels: list[dict], cache_path: Path) -> None:
        """Skip box and segment checks; semantic masks carry no box or segment annotations."""

    def get_labels(self) -> list[dict]:
        """Load semantic labels from cache or scan image-mask paths.

        Returns:
            (list[dict]): List of label dictionaries with mask file paths and image shapes.
        """
        labels = super().get_labels()
        self.mask_files = [lb["mask_file"] for lb in labels]
        return labels

    def load_image(self, i: int, rect_mode: bool = True) -> tuple[np.ndarray, tuple[int, int], tuple[int, int]]:
        """Load an image for semantic segmentation, scaling the short side to imgsz when augmenting with rect_mode."""
        return super().load_image(i, rect_mode=rect_mode, resize_short=self.augment)

    def load_mask(self, index: int, image_shape: tuple[int, int] | None = None) -> np.ndarray:
        """Load a semantic mask and apply optional dataset label mapping.

        Args:
            index (int): Dataset index.
            image_shape (tuple[int, int], optional): Image shape (H, W). Unused here; required by subclasses that
                rasterize masks.

        Returns:
            (np.ndarray): Uint8 mask of class ids at the mask file's native resolution.

        Raises:
            FileNotFoundError: If the mask file is missing or unreadable.
        """
        mask_file = self.labels[index]["mask_file"]
        mode = self.labels[index]["mode"]
        if mode == "P":  # palette PNGs store class ids as indices, not grayscale colors
            with Image.open(mask_file) as im:
                p = np.array(im.getpalette()).reshape(-1, 3)  # gray palettes (e.g. pngquant) hold gray-level class ids
                mask = np.array(im.convert("L") if (p == p[:, :1]).all() else im)
        else:
            mask = cv2.imread(mask_file, cv2.IMREAD_ANYDEPTH)  # grayscale that keeps 16-bit ids
        if mask is None:
            raise FileNotFoundError(f"Semantic mask not found or unreadable: {mask_file}")
        if int(self.data.get("nc", 0)) == 1 and mode == "1":
            mask[mask == 255] = 1  # cv2 expands 1-bit PNG foreground to 255.
        if self.label_mapping:
            mask = self.convert_label(mask, inverse=False)
        return mask.astype(np.uint8, copy=False)

    def convert_label(self, label: np.ndarray, inverse: bool = False) -> np.ndarray:
        """Convert label values using the dataset's label mapping.

        Args:
            label (np.ndarray): Segmentation label array with integer ids in 0-255.
            inverse (bool): If True, apply inverse mapping (mapped -> original).

        Returns:
            (np.ndarray): New uint8 array with converted values.
        """
        lut = self.inverse_lut if inverse else self.label_lut
        return cv2.LUT(label, lut) if label.dtype == np.uint8 else lut[label]  # cv2.LUT needs a uint8 input

    def get_image_and_label(self, index: int) -> dict[str, Any]:
        """Get image, label and semantic mask for the given index.

        Overrides parent to include the semantic mask, served from RAM for the images Mosaic draws from the buffer.

        Args:
            index (int): Dataset index.

        Returns:
            (dict): Label dict with 'img', 'semantic_mask', and metadata.
        """
        label = super().get_image_and_label(index)
        h, w = label["img"].shape[:2]
        mask = self.masks.get(index)
        if mask is None:
            mask = self.load_mask(index, image_shape=(h, w))
            if self.include_class is not None:  # keep only selected classes; remap the rest to the ignore label
                mask[~np.isin(mask, self.include_class)] = 255
            if mask.shape[:2] != (h, w):
                mask = cv2.resize(mask, (w, h), interpolation=cv2.INTER_NEAREST)
            if index in self.buffer:  # image is RAM-resident for mosaic reuse, keep its mask with it
                self.masks[index] = mask
                if len(self.masks) > len(self.buffer):
                    self.masks = {i: self.masks[i] for i in self.buffer if i in self.masks}
        label["semantic_mask"] = mask
        return label


class PolygonSemanticDataset(SemanticDataset, YOLODataset):
    """Semantic segmentation dataset that rasterizes YOLO polygon labels into masks on the fly.

    Used when the dataset YAML lacks 'masks_dir'. Pixels not covered by any polygon become a dedicated background class.
    Requires `add_polygon_background(data)` to be called first: for nc > 1 it bumps `data['nc']` to user_nc + 1 with
    background at `nc - 1`; for nc == 1 it keeps nc=1 and rasterizes a {0=bg, 1=fg} binary mask for use with
    BCEWithLogitsLoss.
    """

    def __init__(self, *args, data: dict, **kwargs):
        """Initialize PolygonSemanticDataset.

        Args:
            *args (Any): Additional positional arguments for the parent class.
            data (dict): Dataset configuration dictionary.
            **kwargs (Any): Additional keyword arguments for the parent class.
        """
        nc = data.get("nc") or len(data.get("names", {}))
        self.bg_class_idx = data.get("bg_class_idx", max(int(nc) - 1, 0))
        super().__init__(*args, data=data, **kwargs)

    # Rebind label scanning to YOLODataset's polygon .txt implementations; the MRO (SemanticDataset, YOLODataset)
    # would otherwise resolve SemanticDataset's PNG-mask hooks and its get_labels, which syncs mask_files from
    # label dicts that polygon labels do not have.
    get_labels = YOLODataset.get_labels
    get_label_files = YOLODataset.get_label_files
    get_cache_hash = YOLODataset.get_cache_hash
    scan_summary = YOLODataset.scan_summary
    verify_args = YOLODataset.verify_args
    result_to_label = YOLODataset.result_to_label
    verify_labels = YOLODataset.verify_labels

    def load_mask(self, index: int, image_shape: tuple[int, int] | None = None) -> np.ndarray:
        """Rasterize this image's polygons into a (H, W) uint8 semantic mask, bg = self.bg_class_idx."""
        h, w = image_shape
        label = self.labels[index]
        cls = label.get("cls")
        segments = label.get("segments") or []
        if cls is None or len(cls) == 0 or len(segments) == 0:
            return np.full((h, w), self.bg_class_idx, dtype=np.uint8)

        # Denormalize polygons (stored as normalized xy) to pixel coordinates at (h, w).
        scale = np.array([w, h], dtype=np.float32)
        polys = [np.asarray(s, dtype=np.float32).reshape(-1, 2) * scale for s in segments]
        # Returns (H, W) instance index map: 0 = no polygon, 1..N = sorted instance index.
        inst, sorted_idx = polygons2masks_overlap((h, w), polys, downsample_ratio=1)
        out = np.full((h, w), self.bg_class_idx, dtype=np.uint8)
        fg = inst > 0
        if int(self.data.get("nc", 0)) == 1:  # binary: fg=1 regardless of label cls value
            out[fg] = 1
        else:
            cls_arr = np.asarray(cls).reshape(-1).astype(np.int32)[sorted_idx]
            out[fg] = cls_arr[inst[fg] - 1].astype(np.uint8)
        return out


class ClassificationDataset:
    """Dataset class for image classification tasks wrapping torchvision ImageFolder functionality.

    This class offers functionalities like image augmentation, caching, and verification. It's designed to efficiently
    handle large datasets for training deep learning models, with optional image transformations and caching mechanisms
    to speed up training.

    Attributes:
        base (torchvision.datasets.ImageFolder): The underlying ImageFolder dataset.
        cache_ram (bool): Indicates if caching in RAM is enabled.
        cache_disk (bool): Indicates if caching on disk is enabled.
        samples (list): A list of lists, each containing the path to an image, its class index, path to its .npy cache
            file, and a None placeholder.
        img_cache (BaseDataset._ImageCache): Contiguous RAM cache of decoded images, set when caching in RAM.
        torch_transforms (callable): PyTorch transforms to be applied to the images.
        root (str): Root directory of the dataset.
        prefix (str): Colored prefix for logging.

    Methods:
        __getitem__: Return transformed image and class index for the given sample index.
        __len__: Return the total number of samples in the dataset.
        verify_images: Verify all images in dataset.
        imread: Read a BGR image, decoding the formats cv2 cannot read through the shared PIL fallback.
        cache_images: Decode images into one contiguous RAM cache.
    """

    def __init__(
        self,
        root: str | Path,
        args: IterableSimpleNamespace,
        augment: bool = False,
        prefix: str = "",
        names: dict[int, str] | None = None,
    ):
        """Initialize YOLO classification dataset with root directory, arguments, augmentations, and cache settings.

        Args:
            root (str | Path): Path to the dataset directory where images are stored in a class-specific folder
                structure.
            args (IterableSimpleNamespace): Configuration containing dataset-related settings such as image size,
                augmentation parameters, and cache settings.
            augment (bool, optional): Whether to apply augmentations to the dataset.
            prefix (str, optional): Split name used as the logging prefix and to select the split's 'fraction'. If
                empty, 'train' is used when augment is True, otherwise 'val'.
            names (dict[int, str], optional): Model class names; class folders are aligned to this order by name and
                folders the model lacks are dropped, since each split's ImageFolder scan is indexed on its own.
        """
        import torchvision  # scope for faster 'import ultralytics'

        # Base class assigned as attribute rather than used as base class to allow for scoping slow torchvision import
        kwargs = {"allow_empty": True} if TORCHVISION_0_18 else {}  # 'allow_empty' first introduced in torchvision 0.18
        self.base = torchvision.datasets.ImageFolder(
            root=root, is_valid_file=lambda x: x.rpartition(".")[-1].lower() in IMG_FORMATS, **kwargs
        )
        is_ndjson = (Path(root).parent / ".ndjson.yaml").is_file()
        self.samples = self.base.samples
        self.root = self.base.root

        # Initialize attributes
        fraction = 1.0 if is_ndjson else get_split_fraction(args.fraction, prefix or ("train" if augment else "val"))
        count = fraction if isinstance(fraction, int) else max(int(fraction > 0), round(len(self.samples) * fraction))
        self.samples = (
            [self.samples[i] for i in np.linspace(0, len(self.samples) - 1, count, dtype=int)]
            if count < len(self.samples)
            else self.samples
        )
        self.prefix = colorstr(f"{prefix}: ") if prefix else ""
        self.cache_ram = args.cache is True or str(args.cache).lower() == "ram"  # cache images into RAM
        self.cache_disk = str(args.cache).lower() == "disk"  # cache images on hard drive as uncompressed *.npy files
        self.samples = self.verify_images()  # filter out bad images
        classes = self.base.classes  # this split's class folders, sorted, indexed by the ImageFolder target
        if args.single_cls:
            index = dict.fromkeys(classes, 0)
        elif is_ndjson:  # folders are the class ids
            index = {c: int(c) for c in classes}
        elif names and not set(classes).isdisjoint(names.values()):  # align to the model's class order by name
            index = {n: i for i, n in names.items()}
        else:  # folder names carry no class meaning, e.g. ImageNet wnids under humanized names
            index = {c: i for i, c in enumerate(classes)}
        extra = {c for c in classes if index.get(c, len(names)) >= len(names)} if names else set()  # not in the model
        n = len(self.samples)
        self.samples = [(f, index[classes[t]]) for f, t in self.samples if classes[t] not in extra]
        if extra:
            LOGGER.warning(
                f"{self.prefix}Skipping {n - len(self.samples)} samples from classes the model lacks: {sorted(extra)}"
            )
        # Same persistent image.npy naming as BaseDataset.npy_files, never rename or relocate existing caches
        self.samples = [[*list(x), Path(x[0]).with_suffix(".npy"), None] for x in self.samples]  # file, index, npy, im
        if self.cache_ram:
            self.cache_images()
        scale = (1.0 - args.scale, 1.0)  # RandomResizedCrop area range, e.g. (0.5, 1.0) for scale=0.5
        self.torch_transforms = (
            classify_augmentations(
                size=args.imgsz,
                scale=scale,
                hflip=args.fliplr,
                vflip=args.flipud,
                erasing=args.erasing,
                auto_augment=args.auto_augment,
                hsv_h=args.hsv_h,
                hsv_s=args.hsv_s,
                hsv_v=args.hsv_v,
            )
            if augment
            else classify_transforms(size=args.imgsz)
        )

    def __getitem__(self, i: int) -> dict:
        """Return transformed image and class index for the given sample index.

        Args:
            i (int): Index of the sample to retrieve.

        Returns:
            (dict): Dictionary containing the image and its class index.
        """
        f, j, fn, im = self.samples[i]  # filename, class index, npy cache path, image
        if self.cache_ram:
            im = self.img_cache[i]
        elif self.cache_disk:
            if not fn.exists() or fn.stat().st_mtime < Path(f).stat().st_mtime:  # missing or stale
                np.save(fn.as_posix(), self.imread(f), allow_pickle=False)
            im = np.load(fn)
        else:  # read image
            im = self.imread(f)  # BGR
        # Convert NumPy array to PIL image
        im = Image.fromarray(cv2.cvtColor(im, cv2.COLOR_BGR2RGB))
        sample = self.torch_transforms(im)
        return {"img": sample, "cls": j}

    def __len__(self) -> int:
        """Return the total number of samples in the dataset."""
        return len(self.samples)

    @staticmethod
    def imread(f: str) -> np.ndarray | None:
        """Read a BGR image with cv2, decoding the formats cv2 cannot read through the shared PIL fallback."""
        return imread(f) if f.lower().endswith(PIL_FALLBACK_SUFFIXES) else imread_unicode(f)

    def cache_images(self) -> None:
        """Decode all images once into a single contiguous uint8 buffer before DataLoader workers fork.

        A Python list of per-image arrays is duplicated into every forked worker by copy-on-write refcounting
        (https://github.com/ultralytics/ultralytics/issues/9824); one shared buffer is read-only across
        workers instead, so RAM stays flat. Original image sizes are preserved for the transforms.
        """
        with ThreadPool(NUM_THREADS) as pool:
            ims = list(
                TQDM(
                    pool.imap(lambda s: self.imread(s[0]), self.samples),
                    total=len(self.samples),
                    desc=f"{self.prefix}Caching images",
                    disable=LOCAL_RANK > 0,
                )
            )
        self.img_cache = BaseDataset._ImageCache(ims)

    def verify_images(self) -> list[tuple]:
        """Verify all images in dataset.

        Returns:
            (list[tuple]): List of valid samples after verification.
        """
        desc = f"{self.prefix}Scanning {self.root}..."
        path = Path(self.root).with_suffix(".cache")  # *.cache file path

        try:
            check_file_speeds([file for (file, _) in self.samples[:5]], prefix=self.prefix)  # check image read speeds
            cache = load_dataset_cache_file(path)  # attempt to load a *.cache file
            assert cache["version"] == DATASET_CACHE_VERSION  # matches current version
            assert cache["hash"] == get_hash([x[0] for x in self.samples] + self.base.classes)  # files and classes
            nf, nc, n, samples = cache.pop("results")  # found, corrupt, total, samples
            if LOCAL_RANK in {-1, 0}:
                d = f"{desc} {nf} images, {nc} corrupt"
                TQDM(None, desc=d, total=n, initial=n)
                if cache["msgs"]:
                    LOGGER.info("\n".join(cache["msgs"]))  # display warnings
            return samples

        except Exception:  # run scan if *.cache retrieval failed, e.g. missing, stale, or truncated
            nf, nc, msgs, samples, x = 0, 0, [], [], {}
            with ThreadPool(NUM_THREADS) as pool:
                results = pool.imap(func=verify_image, iterable=zip(self.samples, repeat(self.prefix)))
                pbar = TQDM(results, desc=desc, total=len(self.samples))
                for sample, nf_f, nc_f, msg in pbar:
                    if nf_f:
                        samples.append(sample)
                    if msg:
                        msgs.append(msg)
                    nf += nf_f
                    nc += nc_f
                    pbar.desc = f"{desc} {nf} images, {nc} corrupt"
                pbar.close()
            if msgs:
                LOGGER.info("\n".join(msgs))
            x["hash"] = get_hash([x[0] for x in self.samples] + self.base.classes)
            x["results"] = nf, nc, len(samples), samples
            x["msgs"] = msgs  # warnings
            save_dataset_cache_file(self.prefix, path, x, DATASET_CACHE_VERSION)
            return samples
