""" Mask R-CNN 实例分割推理与可视化脚本 用法: # 对单张图片进行推理 python predict.py --image dataset/1.jpg --checkpoint checkpoints/mask_rcnn_best.pth # 对整个数据集目录进行推理 python predict.py --dataset dataset --checkpoint checkpoints/mask_rcnn_best.pth # 调整置信度阈值 python predict.py --image dataset/1.jpg --checkpoint checkpoints/mask_rcnn_best.pth --threshold 0.7 # 同时显示 ground truth 对比 python predict.py --image dataset/1.jpg --checkpoint checkpoints/mask_rcnn_best.pth --show-gt """ import argparse import os import sys import json from typing import Dict, Optional import numpy as np from PIL import Image import torch from utils import ( get_label_map, get_image_json_pairs, shapes_to_masks, visualize_instances, ) from dataset import get_val_transforms from model import get_model def load_checkpoint(checkpoint_path: str, device: torch.device): """ 加载训练好的模型检查点。 Args: checkpoint_path: 检查点文件路径 device: 计算设备 Returns: model: 加载好权重的模型 label_map: 标签名→类别ID 映射 """ checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) label_map = checkpoint.get("label_map", {}) num_classes = max(label_map.values()) + 1 # +1 for background model = get_model(num_classes=num_classes, pretrained=False) model.load_state_dict(checkpoint["model_state_dict"]) model = model.to(device) model.eval() # 构建 id→name 映射 id_to_name = {v: k for k, v in label_map.items()} return model, label_map, id_to_name @torch.no_grad() def predict_single( model: torch.nn.Module, image: np.ndarray, device: torch.device, threshold: float = 0.5, ) -> Dict: """ 对单张图像进行推理。 Args: model: 模型 image: [H, W, 3] RGB uint8 图像 device: 计算设备 threshold: 置信度阈值 Returns: 包含 boxes, masks, labels, scores 的字典 """ # 预处理 image_float = image.astype(np.float32) / 255.0 image_tensor = torch.as_tensor(image_float, dtype=torch.float32) image_tensor = image_tensor.permute(2, 0, 1).unsqueeze(0).to(device) # 推理 outputs = model(image_tensor) output = outputs[0] # 过滤低置信度 keep = output["scores"] > threshold boxes = output["boxes"][keep].cpu().numpy() masks = output["masks"][keep].cpu().numpy() labels = output["labels"][keep].cpu().numpy() scores = output["scores"][keep].cpu().numpy() # mask 阈值化(> 0.5) masks = (masks[:, 0, :, :] > 0.5).astype(np.uint8) return { "boxes": boxes, "masks": masks, "labels": labels, "scores": scores, } def visualize_and_save( image: np.ndarray, result: Dict, id_to_name: Dict[int, str], output_path: str, threshold: float = 0.5, ): """ 可视化预测结果并保存图像。 Args: image: 原始图像 [H, W, 3] RGB result: 预测结果字典 id_to_name: 类别ID→名称映射 output_path: 输出路径 threshold: 置信度阈值 """ vis_img = visualize_instances( image, result["boxes"], result["masks"], result["labels"], result["scores"], label_map=id_to_name, score_threshold=threshold, ) os.makedirs(os.path.dirname(output_path), exist_ok=True) Image.fromarray(vis_img).save(output_path) print(f" -> 结果已保存: {output_path}") def visualize_gt( image: np.ndarray, json_path: str, label_map: Dict[str, int], id_to_name: Dict[int, str], output_path: str, ): """可视化 ground truth 标注""" with open(json_path, "r", encoding="utf-8") as f: anno = json.load(f) height = anno.get("imageHeight", image.shape[0]) width = anno.get("imageWidth", image.shape[1]) masks, labels, boxes = shapes_to_masks( anno.get("shapes", []), height, width, label_map ) if len(labels) == 0: print(" -> 无标注实例") return vis_img = visualize_instances( image, boxes, masks, labels, None, id_to_name, 0.0 ) os.makedirs(os.path.dirname(output_path), exist_ok=True) Image.fromarray(vis_img).save(output_path) print(f" -> Ground Truth 已保存: {output_path}") def main(): parser = argparse.ArgumentParser(description="Mask R-CNN 实例分割推理") parser.add_argument("--image", type=str, default=None, help="单张图像路径") parser.add_argument("--dataset", type=str, default=None, help="数据集目录路径(批量推理)") parser.add_argument("--checkpoint", type=str, default="checkpoints/mask_rcnn_best.pth", help="模型检查点路径 (默认: checkpoints/mask_rcnn_best.pth)") parser.add_argument("--output-dir", type=str, default="predictions", help="输出目录 (默认: predictions)") parser.add_argument("--threshold", type=float, default=0.5, help="置信度阈值 (默认: 0.5)") parser.add_argument("--show-gt", action="store_true", help="同时可视化 ground truth 对比") parser.add_argument("--min-size", type=int, default=800, help="输入图像最小边尺寸 (默认: 800)") parser.add_argument("--max-size", type=int, default=1333, help="输入图像最大边尺寸 (默认: 1333)") args = parser.parse_args() # 设置设备 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"使用设备: {device}") # 检查检查点 ckpt_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), args.checkpoint) if not os.path.exists(ckpt_path): print(f"错误: 检查点文件不存在: {ckpt_path}") print("请先运行 train.py 训练模型,或使用 --checkpoint 指定正确的路径。") sys.exit(1) # 加载模型 model, label_map, id_to_name = load_checkpoint(ckpt_path, device) print(f"模型已加载 | 标签映射: {label_map}") print(f"置信度阈值: {args.threshold}") output_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), args.output_dir) os.makedirs(output_dir, exist_ok=True) # 收集要推理的图像 if args.image: image_paths = [args.image] json_paths = [args.image.rsplit(".", 1)[0] + ".json"] elif args.dataset: dataset_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), args.dataset) pairs = get_image_json_pairs(dataset_path) image_paths = [p[0] for p in pairs] json_paths = [p[1] for p in pairs] else: print("错误: 请指定 --image 或 --dataset") sys.exit(1) print(f"\n共 {len(image_paths)} 张图像待推理\n") for idx, (image_path, json_path) in enumerate(zip(image_paths, json_paths)): if not os.path.exists(image_path): print(f"[{idx+1}] 跳过(图像不存在): {image_path}") continue print(f"[{idx+1}/{len(image_paths)}] {os.path.basename(image_path)}") # 读取图像 image = np.array(Image.open(image_path).convert("RGB")) # 推理 result = predict_single(model, image, device, args.threshold) num_detections = len(result["scores"]) print(f" -> 检测到 {num_detections} 个实例") for i in range(num_detections): name = id_to_name.get(int(result["labels"][i]), str(result["labels"][i])) print(f" {name}: {result['scores'][i]:.3f}") # 可视化并保存 base_name = os.path.splitext(os.path.basename(image_path))[0] output_path = os.path.join(output_dir, f"{base_name}_pred.png") visualize_and_save(image, result, id_to_name, output_path, args.threshold) # 可选:可视化 ground truth if args.show_gt and os.path.exists(json_path): gt_output_path = os.path.join(output_dir, f"{base_name}_gt.png") visualize_gt(image, json_path, label_map, id_to_name, gt_output_path) print(f"\n推理完成!结果保存在: {output_dir}") if __name__ == "__main__": main()