from __future__ import annotations import json import random from collections import defaultdict from pathlib import Path from typing import Any import numpy as np from PIL import Image, ImageEnhance import torch from torch import Tensor from torch.utils.data import Dataset from .boxes import box_xyxy_to_cxcywh class CocoDetectionDataset(Dataset): """Minimal COCO detection loader with deterministic category remapping.""" def __init__( self, image_dir: str | Path, annotation_file: str | Path, input_size: int, training: bool, hflip_prob: float = 0.5, scale_range: tuple[float, float] = (0.65, 1.0), mean: tuple[float, float, float] = (0.485, 0.456, 0.406), std: tuple[float, float, float] = (0.229, 0.224, 0.225), ) -> None: self.image_dir = Path(image_dir) self.annotation_file = Path(annotation_file) self.input_size = input_size self.training = training self.hflip_prob = hflip_prob self.scale_range = scale_range self.mean = torch.tensor(mean, dtype=torch.float32)[:, None, None] self.std = torch.tensor(std, dtype=torch.float32)[:, None, None] with self.annotation_file.open("r", encoding="utf-8") as handle: data = json.load(handle) self.images = sorted(data["images"], key=lambda item: item["id"]) annotations: dict[int, list[dict[str, Any]]] = defaultdict(list) for annotation in data["annotations"]: if annotation.get("iscrowd", 0) == 0 and annotation["bbox"][2] > 0 and annotation["bbox"][3] > 0: annotations[annotation["image_id"]].append(annotation) self.annotations = annotations category_ids = sorted(category["id"] for category in data["categories"]) self.category_to_label = {category_id: label for label, category_id in enumerate(category_ids)} self.label_to_category = {label: category_id for category_id, label in self.category_to_label.items()} self.categories = sorted(data["categories"], key=lambda item: self.category_to_label[item["id"]]) def __len__(self) -> int: return len(self.images) def _letterbox( self, image: Image.Image, boxes: Tensor ) -> tuple[Image.Image, Tensor, tuple[float, int, int]]: width, height = image.size scale_jitter = random.uniform(*self.scale_range) if self.training else 1.0 ratio = min(self.input_size / width, self.input_size / height) * scale_jitter resized_width = max(1, round(width * ratio)) resized_height = max(1, round(height * ratio)) image = image.resize((resized_width, resized_height), Image.Resampling.BILINEAR) max_x = self.input_size - resized_width max_y = self.input_size - resized_height if self.training: offset_x = random.randint(0, max_x) if max_x else 0 offset_y = random.randint(0, max_y) if max_y else 0 else: offset_x, offset_y = max_x // 2, max_y // 2 canvas = Image.new("RGB", (self.input_size, self.input_size), (114, 114, 114)) canvas.paste(image, (offset_x, offset_y)) if boxes.numel(): boxes = boxes * ratio boxes[:, [0, 2]] += offset_x boxes[:, [1, 3]] += offset_y return canvas, boxes, (ratio, offset_x, offset_y) def __getitem__(self, index: int) -> tuple[Tensor, dict[str, Tensor]]: image_info = self.images[index] image = Image.open(self.image_dir / image_info["file_name"]).convert("RGB") original_width, original_height = image.size records = self.annotations.get(image_info["id"], []) boxes = [] labels = [] for record in records: x, y, width, height = record["bbox"] boxes.append((x, y, x + width, y + height)) labels.append(self.category_to_label[record["category_id"]]) box_tensor = torch.tensor(boxes, dtype=torch.float32).reshape(-1, 4) label_tensor = torch.tensor(labels, dtype=torch.int64) if self.training and random.random() < self.hflip_prob: image = image.transpose(Image.Transpose.FLIP_LEFT_RIGHT) if box_tensor.numel(): old_x0 = box_tensor[:, 0].clone() box_tensor[:, 0] = original_width - box_tensor[:, 2] box_tensor[:, 2] = original_width - old_x0 if self.training: image = ImageEnhance.Color(image).enhance(random.uniform(0.8, 1.2)) image = ImageEnhance.Contrast(image).enhance(random.uniform(0.8, 1.2)) image, box_tensor, (ratio, offset_x, offset_y) = self._letterbox(image, box_tensor) if box_tensor.numel(): box_tensor = box_xyxy_to_cxcywh(box_tensor) / self.input_size valid = (box_tensor[:, 2] > 1e-4) & (box_tensor[:, 3] > 1e-4) box_tensor = box_tensor[valid].clamp(0.0, 1.0) label_tensor = label_tensor[valid] image_array = np.asarray(image, dtype=np.float32).copy() / 255.0 image_tensor = torch.from_numpy(image_array).permute(2, 0, 1) image_tensor = (image_tensor - self.mean) / self.std target = { "boxes": box_tensor, "labels": label_tensor, "image_id": torch.tensor(image_info["id"], dtype=torch.int64), "original_size": torch.tensor([original_height, original_width], dtype=torch.int64), "transform": torch.tensor([ratio, offset_x, offset_y], dtype=torch.float32), } return image_tensor, target def detection_collate(batch): images, targets = zip(*batch, strict=True) return torch.stack(images), list(targets) def build_dataset(config: dict, root: str | Path, split: str) -> CocoDetectionDataset: data = config["data"] root = Path(root) training = split == "train" return CocoDetectionDataset( root / data[f"{split}_image_dir"], root / data[f"{split}_annotations"], input_size=int(config["model"]["input_size"]), training=training, hflip_prob=float(data.get("hflip_prob", 0.5)), scale_range=tuple(data.get("scale_range", (0.65, 1.0))) if training else (1.0, 1.0), mean=tuple(data.get("mean", (0.485, 0.456, 0.406))), std=tuple(data.get("std", (0.229, 0.224, 0.225))), )