""" 自定义 Dataset 类 支持 LabelMe 标注格式(polygon / linestrip / rectangle / circle) 内置数据增强(水平翻转、亮度调整) """ import os import json from typing import List, Dict, Tuple, Optional import numpy as np from PIL import Image import torch from torch.utils.data import Dataset from utils import ( get_label_map, shapes_to_masks, get_image_json_pairs, split_dataset, ) class LabelMeDataset(Dataset): """ LabelMe 格式的实例分割数据集。 数据目录结构: dataset/ ├── 1.jpg ├── 1.json ├── 2.jpg ├── 2.json └── ... 每个 JSON 文件包含 LabelMe 标注,shapes 列表中每个 shape 有: - label: 类别名 - points: [[x, y], ...] 点坐标 - shape_type: "polygon" / "linestrip" / "rectangle" / "circle" linestrip 类型会自动将相同 label 的多条线段合并为一个多边形实例。 """ def __init__( self, dataset_dir: str, label_map: Optional[Dict[str, int]] = None, pairs: Optional[List[Tuple[str, str]]] = None, transforms: Optional[object] = None, ): """ Args: dataset_dir: 数据集目录路径 label_map: 标签名→类别ID 映射(若为 None 则自动构建) pairs: 指定使用的 (image_path, json_path) 配对(若为 None 则扫描全部) transforms: 数据增强变换 """ self.dataset_dir = dataset_dir self.transforms = transforms # 自动构建或使用传入的标签映射 if label_map is None: label_map = get_label_map(dataset_dir) self.label_map = label_map # 获取图像-标注配对 if pairs is None: pairs = get_image_json_pairs(dataset_dir) self.pairs = pairs def __len__(self) -> int: return len(self.pairs) def __getitem__(self, idx: int) -> Tuple[torch.Tensor, Dict]: image_path, json_path = self.pairs[idx] # 读取图像 image = Image.open(image_path).convert("RGB") width, height = image.size # 读取标注 with open(json_path, "r", encoding="utf-8") as f: anno = json.load(f) # 确保 imageHeight / imageWidth 与实际一致 anno_h = anno.get("imageHeight", height) anno_w = anno.get("imageWidth", width) # 转换 shapes → masks, labels, boxes masks, labels, boxes = shapes_to_masks( anno.get("shapes", []), anno_h, anno_w, self.label_map, ) # 转为 numpy 数组 image_np = np.array(image, dtype=np.float32) / 255.0 # [H, W, 3] 0~1 # 构建 target 字典 num_instances = len(labels) target = { "boxes": torch.as_tensor(boxes, dtype=torch.float32), "labels": torch.as_tensor(labels, dtype=torch.int64), "masks": torch.as_tensor(masks, dtype=torch.uint8), "image_id": torch.tensor([idx]), "area": torch.zeros((num_instances,), dtype=torch.float32), "iscrowd": torch.zeros((num_instances,), dtype=torch.int64), } # 计算 area if num_instances > 0: target["area"] = (target["masks"] > 0).sum(dim=(1, 2)).float() # 应用数据增强 if self.transforms is not None: image_np, target = self.transforms(image_np, target) # 转换为 tensor [C, H, W] image_tensor = torch.as_tensor(image_np, dtype=torch.float32) if image_tensor.dim() == 3: image_tensor = image_tensor.permute(2, 0, 1) return image_tensor, target def get_label_map(self) -> Dict[str, int]: return self.label_map def get_num_classes(self) -> int: """返回类别数(不含背景)""" return max(self.label_map.values()) if self.label_map else 0 # ============================================================ # 数据增强 # ============================================================ class Compose: """组合多个变换""" def __init__(self, transforms: List): self.transforms = transforms def __call__(self, image: np.ndarray, target: dict) -> Tuple[np.ndarray, dict]: for t in self.transforms: image, target = t(image, target) return image, target class RandomHorizontalFlip: """随机水平翻转(同时翻转图像、mask、boxes)""" def __init__(self, prob: float = 0.5): self.prob = prob def __call__(self, image: np.ndarray, target: dict) -> Tuple[np.ndarray, dict]: if np.random.random() < self.prob: image = np.ascontiguousarray(image[:, ::-1, :]) _, width = image.shape[:2] # 翻转 mask if "masks" in target and len(target["masks"]) > 0: target["masks"] = torch.flip(target["masks"], dims=[2]) target["masks"] = target["masks"].contiguous() # 翻转 boxes if "boxes" in target and len(target["boxes"]) > 0: boxes = target["boxes"].clone() boxes[:, 0] = width - target["boxes"][:, 2] boxes[:, 2] = width - target["boxes"][:, 0] target["boxes"] = boxes return image, target class RandomBrightness: """随机亮度调整""" def __init__(self, brightness_range: float = 0.2): self.brightness_range = brightness_range def __call__(self, image: np.ndarray, target: dict) -> Tuple[np.ndarray, dict]: factor = 1.0 + np.random.uniform( -self.brightness_range, self.brightness_range ) image = np.clip(image * factor, 0.0, 1.0) return image, target class RandomContrast: """随机对比度调整""" def __init__(self, contrast_range: float = 0.2): self.contrast_range = contrast_range def __call__(self, image: np.ndarray, target: dict) -> Tuple[np.ndarray, dict]: factor = 1.0 + np.random.uniform( -self.contrast_range, self.contrast_range ) mean = image.mean() image = np.clip((image - mean) * factor + mean, 0.0, 1.0) return image, target class Normalize: """标准化(使用 ImageNet 均值和标准差)""" MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32) STD = np.array([0.229, 0.224, 0.225], dtype=np.float32) def __call__(self, image: np.ndarray, target: dict) -> Tuple[np.ndarray, dict]: image = (image - self.MEAN) / self.STD return image, target def get_train_transforms() -> Compose: """获取训练集数据增强""" return Compose([ RandomHorizontalFlip(prob=0.5), RandomBrightness(brightness_range=0.2), RandomContrast(contrast_range=0.2), ]) def get_val_transforms() -> Compose: """获取验证集数据增强(仅标准化,不做随机变换)""" return Compose([])