Files
segmentation/utils.py
T

465 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
工具函数模块
包含: 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 二值 mask0 或 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 二值 mask0 或 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