""" 工具函数模块 包含: LabelMe 标注解析、linestrip 合并为多边形、多边形转 mask、可视化等 """ import json import os import math import random from typing import List, Dict, Tuple, Optional import numpy as np from PIL import Image, ImageDraw # ============================================================ # 1. 标签映射 # ============================================================ def get_label_map(dataset_dir: str) -> Dict[str, int]: """ 扫描数据集中所有 LabelMe JSON 文件,自动构建标签名→类别ID 映射。 类别从 1 开始(0 保留给背景)。 Args: dataset_dir: 数据集目录路径 Returns: 标签名 → 类别ID 的字典,如 {"crosswalk": 1} """ labels = set() for fname in sorted(os.listdir(dataset_dir)): if not fname.endswith(".json"): continue with open(os.path.join(dataset_dir, fname), "r", encoding="utf-8") as f: data = json.load(f) for shape in data.get("shapes", []): labels.add(shape["label"]) # 排序保证可复现 label_map = {label: idx + 1 for idx, label in enumerate(sorted(labels))} return label_map # ============================================================ # 2. Linestrip 合并为多边形 # ============================================================ def _point_dist(p1, p2) -> float: """计算两点之间的欧氏距离""" return math.sqrt((p1[0] - p2[0]) ** 2 + (p1[1] - p2[1]) ** 2) def merge_linestrips_to_polygon(strips: List[List[List[float]]]) -> List[List[float]]: """ 将多条 linestrip 合并为一个闭合多边形。 策略:以第一条 strip 为起点,每次找到与当前多边形末端点距离最近的 下一条 strip 的端点,将该 strip(可能反转)拼接到多边形末尾, 最后闭合多边形。 对于单条 linestrip,直接闭合。 Args: strips: 多条 linestrip 的点列表,每条 strip 是 [[x, y], ...] Returns: 闭合多边形的点列表 [[x, y], ..., [x, y]](首尾相同) """ if len(strips) == 0: return [] # 深拷贝避免修改原始数据 strips = [list(s) for s in strips] if len(strips) == 1: polygon = list(strips[0]) polygon.append(list(polygon[0])) # 闭合 return polygon # 从第一条 strip 开始 polygon = list(strips[0]) remaining = strips[1:] while remaining: last_point = polygon[-1] best_idx = 0 best_dist = float("inf") best_reverse = False for i, strip in enumerate(remaining): # 尝试正向连接(strip 的起点与 polygon 末端最近) d_start = _point_dist(last_point, strip[0]) # 尝试反向连接(strip 的终点与 polygon 末端最近) d_end = _point_dist(last_point, strip[-1]) if d_start < best_dist: best_dist = d_start best_idx = i best_reverse = False if d_end < best_dist: best_dist = d_end best_idx = i best_reverse = True strip = remaining.pop(best_idx) if best_reverse: polygon.extend(list(reversed(strip))) else: polygon.extend(strip) # 闭合多边形 polygon.append(list(polygon[0])) return polygon # ============================================================ # 3. 多边形 / Linestrip → Mask # ============================================================ def shapes_to_masks( shapes: List[dict], height: int, width: int, label_map: Dict[str, int], ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ 将 LabelMe 的 shapes 转换为实例分割所需的 masks、labels、boxes。 处理逻辑: - polygon 类型:直接作为多边形填充 - linestrip 类型:将相同 label 的 linestrip 合并为一个多边形后填充 - rectangle / circle 类型:转换为多边形后处理 Args: shapes: LabelMe JSON 中的 shapes 列表 height: 图像高度 width: 图像宽度 label_map: 标签名→类别ID 映射 Returns: masks: [N, H, W] uint8 二值 mask(0 或 1) labels: [N] int64 类别ID boxes: [N, 4] float32 格式 (x1, y1, x2, y2) """ # 按 label 分组 linestrip(同 label 的 linestrip 合并为一个实例) linestrip_groups: Dict[str, List[List]] = {} # 非 linestrip 的 shape 逐个处理 other_shapes: List[dict] = [] for shape in shapes: shape_type = shape.get("shape_type", "polygon") if shape_type == "linestrip": label = shape["label"] linestrip_groups.setdefault(label, []).append(shape["points"]) else: other_shapes.append(shape) all_masks = [] all_labels = [] # 处理合并后的 linestrip for label, strips in linestrip_groups.items(): polygon = merge_linestrips_to_polygon(strips) if len(polygon) < 3: continue mask = _polygon_to_mask(polygon, height, width) if mask.sum() == 0: continue all_masks.append(mask) all_labels.append(label_map.get(label, 1)) # 处理其他 shape 类型 for shape in other_shapes: label = shape["label"] shape_type = shape.get("shape_type", "polygon") points = shape["points"] if shape_type == "polygon": polygon = points elif shape_type == "rectangle": # 矩形转多边形 x1, y1 = points[0] x2, y2 = points[1] polygon = [[x1, y1], [x2, y1], [x2, y2], [x1, y2]] elif shape_type == "circle": # 圆形转多边形(近似) cx, cy = points[0] rx = points[1][0] - cx ry = points[1][1] - cy polygon = [] for angle in range(0, 360, 5): rad = math.radians(angle) polygon.append([cx + rx * math.cos(rad), cy + ry * math.sin(rad)]) else: continue if len(polygon) < 3: continue mask = _polygon_to_mask(polygon, height, width) if mask.sum() == 0: continue all_masks.append(mask) all_labels.append(label_map.get(label, 1)) if len(all_masks) == 0: return ( np.zeros((0, height, width), dtype=np.uint8), np.zeros((0,), dtype=np.int64), np.zeros((0, 4), dtype=np.float32), ) masks = np.stack(all_masks, axis=0) # [N, H, W] labels = np.array(all_labels, dtype=np.int64) # 从 mask 计算包围框 boxes = masks_to_boxes(masks) return masks, labels, boxes def _polygon_to_mask( polygon: List[List[float]], height: int, width: int, ) -> np.ndarray: """ 将多边形转换为二值 mask。 Args: polygon: 多边形点列表 [[x, y], ...] height: 图像高度 width: 图像宽度 Returns: [H, W] uint8 二值 mask(0 或 1) """ img = Image.new("L", (width, height), 0) ImageDraw.Draw(img).polygon( [(float(p[0]), float(p[1])) for p in polygon], fill=1, outline=1, ) return np.array(img, dtype=np.uint8) def masks_to_boxes(masks: np.ndarray) -> np.ndarray: """ 从一组二值 mask 计算每个实例的包围框。 Args: masks: [N, H, W] uint8 Returns: [N, 4] float32, 格式 (x1, y1, x2, y2) """ boxes = np.zeros((masks.shape[0], 4), dtype=np.float32) for i, mask in enumerate(masks): ys, xs = np.where(mask > 0) if len(xs) == 0: continue boxes[i] = [xs.min(), ys.min(), xs.max(), ys.max()] return boxes # ============================================================ # 4. 可视化 # ============================================================ def visualize_instances( image: np.ndarray, boxes: np.ndarray, masks: np.ndarray, labels: np.ndarray, scores: Optional[np.ndarray] = None, label_map: Optional[Dict[int, str]] = None, score_threshold: float = 0.5, ) -> np.ndarray: """ 在图像上绘制实例分割结果(包围框 + mask 叠加)。 Args: image: [H, W, 3] uint8 BGR 或 RGB 图像 boxes: [N, 4] 包围框 masks: [N, H, W] 二值 mask labels: [N] 类别ID scores: [N] 置信度(可选) label_map: 类别ID→名称映射(可选) score_threshold: 置信度阈值 Returns: 绘制后的图像 [H, W, 3] uint8 """ try: import cv2 except ImportError: # 如果没有 cv2,退化为 PIL 绘制 return _visualize_instances_pil( image, boxes, masks, labels, scores, label_map, score_threshold ) img = image.copy() if img.dtype != np.uint8: img = (img * 255).astype(np.uint8) # 生成固定颜色 np.random.seed(42) num_classes = max(int(labels.max()) + 1, 10) if len(labels) > 0 else 10 colors = np.random.randint(0, 255, size=(num_classes, 3), dtype=np.uint8) for i in range(len(boxes)): if scores is not None and scores[i] < score_threshold: continue color = colors[int(labels[i])].tolist() x1, y1, x2, y2 = boxes[i].astype(int) # 绘制 mask 半透明叠加 mask = masks[i] colored_mask = np.zeros_like(img) for c in range(3): colored_mask[:, :, c] = mask * color[c] img = cv2.addWeighted(img, 1.0, colored_mask, 0.5, 0) # 绘制包围框 cv2.rectangle(img, (x1, y1), (x2, y2), color, 2) # 绘制标签 label_text = label_map.get(int(labels[i]), str(labels[i])) if label_map else str(labels[i]) if scores is not None: label_text = f"{label_text}: {scores[i]:.2f}" (tw, th), _ = cv2.getTextSize(label_text, cv2.FONT_HERSHEY_SIMPLEX, 0.6, 1) cv2.rectangle(img, (x1, y1 - th - 6), (x1 + tw + 4, y1), color, -1) cv2.putText( img, label_text, (x1 + 2, y1 - 4), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (255, 255, 255), 1, cv2.LINE_AA, ) return img def _visualize_instances_pil( image: np.ndarray, boxes: np.ndarray, masks: np.ndarray, labels: np.ndarray, scores: Optional[np.ndarray], label_map: Optional[Dict[int, str]], score_threshold: float, ) -> np.ndarray: """使用 PIL 的可视化后备方案(当 cv2 不可用时)""" from PIL import Image as PILImage, ImageDraw as PILDraw, ImageFont if image.dtype != np.uint8: image = (image * 255).astype(np.uint8) # 转 RGB if image.shape[2] == 3: pil_img = PILImage.fromarray(image[..., ::-1]) if image.dtype == np.uint8 else PILImage.fromarray(image) else: pil_img = PILImage.fromarray(image) np.random.seed(42) num_classes = max(int(labels.max()) + 1, 10) if len(labels) > 0 else 10 colors = np.random.randint(0, 255, size=(num_classes, 3), dtype=np.uint8) overlay = pil_img.copy() draw_overlay = PILDraw.Draw(overlay) draw = PILDraw.Draw(pil_img) for i in range(len(boxes)): if scores is not None and scores[i] < score_threshold: continue color = tuple(colors[int(labels[i])].tolist()) mask = masks[i] # 用颜色填充 mask 区域 mask_img = PILImage.new("RGB", pil_img.size, (0, 0, 0)) mask_draw = PILDraw.Draw(mask_img) mask_pil = PILImage.fromarray((mask * 255).astype(np.uint8)) # 叠加半透明 mask for c in range(3): channel = np.array(pil_img)[:, :, c].astype(float) channel[mask > 0] = channel[mask > 0] * 0.5 + color[c] * 0.5 pil_img_array = np.array(pil_img) pil_img_array[:, :, c] = channel pil_img = PILImage.fromarray(pil_img_array.astype(np.uint8)) draw = PILDraw.Draw(pil_img) for i in range(len(boxes)): if scores is not None and scores[i] < score_threshold: continue color = tuple(colors[int(labels[i])].tolist()) x1, y1, x2, y2 = boxes[i].astype(int) draw.rectangle([x1, y1, x2, y2], outline=color, width=2) label_text = label_map.get(int(labels[i]), str(labels[i])) if label_map else str(labels[i]) if scores is not None: label_text = f"{label_text}: {scores[i]:.2f}" draw.text((x1 + 2, max(0, y1 - 15)), label_text, fill=color) return np.array(pil_img) # ============================================================ # 5. 数据集划分 # ============================================================ def get_image_json_pairs(dataset_dir: str) -> List[Tuple[str, str]]: """ 扫描数据集目录,返回 (image_path, json_path) 配对列表。 Args: dataset_dir: 数据集目录 Returns: [(image_path, json_path), ...] """ pairs = [] for fname in sorted(os.listdir(dataset_dir)): if not fname.endswith(".json"): continue json_path = os.path.join(dataset_dir, fname) with open(json_path, "r", encoding="utf-8") as f: data = json.load(f) image_name = data.get("imagePath", fname.replace(".json", ".jpg")) image_path = os.path.join(dataset_dir, image_name) if not os.path.exists(image_path): # 尝试其他常见扩展名 for ext in [".jpg", ".jpeg", ".png", ".bmp"]: alt_path = os.path.join(dataset_dir, fname.replace(".json", ext)) if os.path.exists(alt_path): image_path = alt_path break pairs.append((image_path, json_path)) return pairs def split_dataset( pairs: List[Tuple[str, str]], val_ratio: float = 0.2, seed: int = 42, ) -> Tuple[List, List]: """ 将数据集划分为训练集和验证集。 Args: pairs: (image_path, json_path) 配对列表 val_ratio: 验证集比例 seed: 随机种子 Returns: train_pairs, val_pairs """ rng = random.Random(seed) indices = list(range(len(pairs))) rng.shuffle(indices) val_size = max(1, int(len(pairs) * val_ratio)) if len(pairs) > 1 else 0 val_indices = set(indices[:val_size]) train_pairs = [pairs[i] for i in indices if i not in val_indices] val_pairs = [pairs[i] for i in val_indices] return train_pairs, val_pairs