""" Mask R-CNN 实例分割训练脚本 用法: # 使用默认参数训练 python train.py # 自定义参数训练 python train.py --dataset dataset --epochs 50 --batch-size 2 --lr 0.005 # 冻结骨干网络(小数据集推荐) python train.py --freeze-backbone --epochs 100 # 从检查点恢复训练 python train.py --resume checkpoints/mask_rcnn_epoch_20.pth 功能: 1. 自动扫描 dataset 目录中的 LabelMe 标注 2. 支持 polygon / linestrip / rectangle / circle 标注类型 3. 自动划分训练集/验证集 4. 支持 COCO 预训练权重微调 5. 支持冻结骨干网络(减少过拟合) 6. 训练过程中保存检查点和最佳模型 7. 记录训练日志(损失曲线) """ import argparse import os import sys import time import json from typing import Dict, List import numpy as np import torch from torch.utils.data import DataLoader from utils import get_label_map, get_image_json_pairs, split_dataset, visualize_instances from dataset import LabelMeDataset, get_train_transforms, get_val_transforms from model import get_model, freeze_backbone, get_optimizer, get_lr_scheduler def collate_fn(batch): """ 自定义 collate 函数:Mask R-CNN 的每个样本可能有不同数量的实例, 不能使用默认的 stack 方式,需要返回 list。 """ return tuple(zip(*batch)) def train_one_epoch( model: torch.nn.Module, optimizer: torch.optim.Optimizer, data_loader: DataLoader, device: torch.device, epoch: int, log_interval: int = 10, ) -> Dict[str, float]: """ 训练一个 epoch。 Args: model: 模型 optimizer: 优化器 data_loader: 训练数据加载器 device: 计算设备 (cuda / cpu) epoch: 当前 epoch 编号 log_interval: 日志打印间隔 Returns: 各项损失的平均值 """ model.train() total_loss = 0.0 loss_components = { "loss_classifier": 0.0, "loss_box_reg": 0.0, "loss_mask": 0.0, "loss_objectness": 0.0, "loss_rpn_box_reg": 0.0, } num_batches = 0 try: from tqdm import tqdm pbar = tqdm(data_loader, desc=f"Epoch {epoch}") except ImportError: pbar = data_loader for batch_idx, (images, targets) in enumerate(pbar): # 过滤无实例的样本 valid_images = [] valid_targets = [] for img, tgt in zip(images, targets): if len(tgt["boxes"]) > 0: valid_images.append(img) valid_targets.append(tgt) if len(valid_images) == 0: continue images = [img.to(device) for img in valid_images] targets = [{k: v.to(device) for k, v in t.items()} for t in valid_targets] # 前向传播(训练模式下返回 loss dict) loss_dict = model(images, targets) losses = sum(loss for loss in loss_dict.values()) # 反向传播 optimizer.zero_grad() losses.backward() optimizer.step() # 记录损失 total_loss += losses.item() num_batches += 1 for k, v in loss_dict.items(): loss_components[k] += v.item() # 更新进度条 if hasattr(pbar, 'set_postfix'): pbar.set_postfix({"loss": f"{losses.item():.4f}"}) avg_loss = total_loss / max(num_batches, 1) for k in loss_components: loss_components[k] /= max(num_batches, 1) loss_components["total"] = avg_loss return loss_components @torch.no_grad() def evaluate( model: torch.nn.Module, data_loader: DataLoader, device: torch.device, ) -> float: """ 验证集评估:计算平均检测置信度作为近似指标。 (完整的 COCO mAP 评估需要 pycocotools,这里使用简化指标避免额外依赖) Args: model: 模型 data_loader: 验证数据加载器 device: 计算设备 Returns: 平均置信度(0~1) """ model.eval() total_score = 0.0 total_instances = 0 for images, targets in data_loader: images = [img.to(device) for img in images] outputs = model(images) for output in outputs: scores = output["scores"] if len(scores) > 0: # 取置信度 > 0.5 的预测 high_conf = scores[scores > 0.5] if len(high_conf) > 0: total_score += high_conf.mean().item() total_instances += 1 return total_score / max(total_instances, 1) def save_checkpoint( model: torch.nn.Module, optimizer: torch.optim.Optimizer, scheduler, epoch: int, label_map: Dict[str, int], loss_history: List[Dict], path: str, ): """保存训练检查点""" torch.save({ "epoch": epoch, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "scheduler_state_dict": scheduler.state_dict() if scheduler else None, "label_map": label_map, "loss_history": loss_history, }, path) def main(): parser = argparse.ArgumentParser(description="Mask R-CNN 实例分割训练") parser.add_argument("--dataset", type=str, default="dataset", help="数据集目录路径 (默认: dataset)") parser.add_argument("--output-dir", type=str, default="checkpoints", help="检查点输出目录 (默认: checkpoints)") parser.add_argument("--epochs", type=int, default=50, help="训练轮数 (默认: 50)") parser.add_argument("--batch-size", type=int, default=2, help="批大小 (默认: 2)") parser.add_argument("--lr", type=float, default=0.005, help="学习率 (默认: 0.005)") parser.add_argument("--momentum", type=float, default=0.9, help="SGD 动量 (默认: 0.9)") parser.add_argument("--weight-decay", type=float, default=0.0005, help="权重衰减 (默认: 0.0005)") parser.add_argument("--lr-step-size", type=int, default=10, help="学习率衰减间隔 (默认: 10)") parser.add_argument("--lr-gamma", type=float, default=0.1, help="学习率衰减系数 (默认: 0.1)") parser.add_argument("--val-ratio", type=float, default=0.2, help="验证集比例 (默认: 0.2)") parser.add_argument("--num-workers", type=int, default=0, help="数据加载线程数 (Windows 建议 0) (默认: 0)") parser.add_argument("--no-pretrained", action="store_true", help="不使用 COCO 预训练权重") parser.add_argument("--freeze-backbone", action="store_true", help="冻结骨干网络(小数据集推荐)") parser.add_argument("--resume", type=str, default=None, help="从检查点恢复训练的路径") parser.add_argument("--save-freq", type=int, default=5, help="每多少 epoch 保存一次检查点 (默认: 5)") parser.add_argument("--min-size", type=int, default=800, help="输入图像最小边尺寸 (默认: 800)") parser.add_argument("--max-size", type=int, default=1333, help="输入图像最大边尺寸 (默认: 1333)") parser.add_argument("--seed", type=int, default=42, help="随机种子 (默认: 42)") args = parser.parse_args() # 设置随机种子 torch.manual_seed(args.seed) np.random.seed(args.seed) # 设置设备 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"使用设备: {device}") # 构建标签映射 dataset_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), args.dataset) label_map = get_label_map(dataset_path) print(f"标签映射: {label_map}") num_classes = max(label_map.values()) + 1 # +1 for background print(f"类别数(含背景): {num_classes}") # 获取数据配对并划分 all_pairs = get_image_json_pairs(dataset_path) print(f"数据集总样本数: {len(all_pairs)}") train_pairs, val_pairs = split_dataset(all_pairs, val_ratio=args.val_ratio, seed=args.seed) print(f"训练集: {len(train_pairs)} | 验证集: {len(val_pairs)}") if len(train_pairs) == 0: print("错误: 训练集为空!请检查数据集目录。") sys.exit(1) # 创建数据集和数据加载器 train_dataset = LabelMeDataset( dataset_dir=dataset_path, label_map=label_map, pairs=train_pairs, transforms=get_train_transforms(), ) val_dataset = LabelMeDataset( dataset_dir=dataset_path, label_map=label_map, pairs=val_pairs, transforms=get_val_transforms(), ) if len(val_pairs) > 0 else None train_loader = DataLoader( train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers, collate_fn=collate_fn, drop_last=False, ) val_loader = DataLoader( val_dataset, batch_size=1, shuffle=False, num_workers=args.num_workers, collate_fn=collate_fn, ) if val_dataset else None # 创建模型 pretrained = not args.no_pretrained model = get_model( num_classes=num_classes, pretrained=pretrained, min_size=args.min_size, max_size=args.max_size, ) print(f"预训练权重: {'COCO' if pretrained else '无'}") if args.freeze_backbone: model = freeze_backbone(model, freeze=True) num_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) num_total = sum(p.numel() for p in model.parameters()) print(f"骨干网络已冻结 | 可训练参数: {num_trainable:,} / {num_total:,}") else: num_total = sum(p.numel() for p in model.parameters()) print(f"可训练参数: {num_total:,}") model = model.to(device) # 创建优化器和调度器 optimizer = get_optimizer(model, lr=args.lr, momentum=args.momentum, weight_decay=args.weight_decay) scheduler = get_lr_scheduler(optimizer, step_size=args.lr_step_size, gamma=args.lr_gamma) # 恢复训练 start_epoch = 0 loss_history = [] best_score = 0.0 if args.resume and os.path.exists(args.resume): checkpoint = torch.load(args.resume, map_location=device, weights_only=False) model.load_state_dict(checkpoint["model_state_dict"]) optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) if checkpoint.get("scheduler_state_dict") and scheduler: scheduler.load_state_dict(checkpoint["scheduler_state_dict"]) start_epoch = checkpoint["epoch"] + 1 loss_history = checkpoint.get("loss_history", []) best_score = max((h.get("val_score", 0) for h in loss_history), default=0.0) print(f"从 epoch {start_epoch} 恢复训练 (best_score={best_score:.4f})") # 创建输出目录 output_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), args.output_dir) os.makedirs(output_dir, exist_ok=True) # 保存标签映射 with open(os.path.join(output_dir, "label_map.json"), "w", encoding="utf-8") as f: json.dump(label_map, f, ensure_ascii=False, indent=2) # 训练循环 print(f"\n{'='*60}") print(f"开始训练 | 总轮数: {args.epochs} | 批大小: {args.batch_size} | 学习率: {args.lr}") print(f"{'='*60}\n") for epoch in range(start_epoch, args.epochs): epoch_start = time.time() # 训练 train_metrics = train_one_epoch( model, optimizer, train_loader, device, epoch + 1 ) # 验证 val_score = 0.0 if val_loader is not None: val_score = evaluate(model, val_loader, device) # 更新学习率 scheduler.step() epoch_time = time.time() - epoch_start current_lr = optimizer.param_groups[0]["lr"] # 记录损失历史 epoch_record = { "epoch": epoch + 1, "train_loss": train_metrics["total"], "loss_classifier": train_metrics["loss_classifier"], "loss_box_reg": train_metrics["loss_box_reg"], "loss_mask": train_metrics["loss_mask"], "loss_objectness": train_metrics["loss_objectness"], "loss_rpn_box_reg": train_metrics["loss_rpn_box_reg"], "val_score": val_score, "lr": current_lr, "time": round(epoch_time, 1), } loss_history.append(epoch_record) # 打印日志 print( f"Epoch {epoch+1}/{args.epochs} | " f"Loss: {train_metrics['total']:.4f} " f"(cls:{train_metrics['loss_classifier']:.4f} " f"box:{train_metrics['loss_box_reg']:.4f} " f"mask:{train_metrics['loss_mask']:.4f}) | " f"Val: {val_score:.4f} | " f"LR: {current_lr:.6f} | " f"Time: {epoch_time:.1f}s" ) # 保存最佳模型 if val_score > best_score: best_score = val_score best_path = os.path.join(output_dir, "mask_rcnn_best.pth") save_checkpoint(model, optimizer, scheduler, epoch + 1, label_map, loss_history, best_path) print(f" -> 新最佳模型已保存 (score={best_score:.4f})") # 定期保存检查点 if (epoch + 1) % args.save_freq == 0 or epoch + 1 == args.epochs: ckpt_path = os.path.join(output_dir, f"mask_rcnn_epoch_{epoch+1}.pth") save_checkpoint(model, optimizer, scheduler, epoch + 1, label_map, loss_history, ckpt_path) print(f" -> 检查点已保存: {ckpt_path}") # 保存最终模型 final_path = os.path.join(output_dir, "mask_rcnn_final.pth") save_checkpoint(model, optimizer, scheduler, args.epochs, label_map, loss_history, final_path) print(f"\n训练完成!最终模型: {final_path}") # 保存损失历史到 JSON with open(os.path.join(output_dir, "loss_history.json"), "w", encoding="utf-8") as f: json.dump(loss_history, f, ensure_ascii=False, indent=2) # 绘制损失曲线 try: _plot_loss_curve(loss_history, output_dir) except Exception as e: print(f"警告: 无法绘制损失曲线 ({e})") print(f"最佳验证分数: {best_score:.4f}") def _plot_loss_curve(loss_history: List[Dict], output_dir: str): """绘制并保存损失曲线图""" import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt epochs = [h["epoch"] for h in loss_history] total_loss = [h["train_loss"] for h in loss_history] cls_loss = [h["loss_classifier"] for h in loss_history] box_loss = [h["loss_box_reg"] for h in loss_history] mask_loss = [h["loss_mask"] for h in loss_history] val_scores = [h.get("val_score", 0) for h in loss_history] fig, axes = plt.subplots(1, 2, figsize=(14, 5)) # 损失曲线 ax1 = axes[0] ax1.plot(epochs, total_loss, "b-", label="Total Loss", linewidth=2) ax1.plot(epochs, cls_loss, "r--", label="Classifier Loss") ax1.plot(epochs, box_loss, "g--", label="Box Reg Loss") ax1.plot(epochs, mask_loss, "m--", label="Mask Loss") ax1.set_xlabel("Epoch") ax1.set_ylabel("Loss") ax1.set_title("Training Loss") ax1.legend() ax1.grid(True, alpha=0.3) # 验证分数曲线 ax2 = axes[1] ax2.plot(epochs, val_scores, "b-o", label="Val Score") ax2.set_xlabel("Epoch") ax2.set_ylabel("Score") ax2.set_title("Validation Score") ax2.legend() ax2.grid(True, alpha=0.3) plt.tight_layout() curve_path = os.path.join(output_dir, "loss_curve.png") plt.savefig(curve_path, dpi=150, bbox_inches="tight") plt.close() print(f"损失曲线图已保存: {curve_path}") if __name__ == "__main__": main() """ # 训练(小数据集推荐冻结骨干网络) python train.py --freeze-backbone --epochs 100 # 推理 python predict.py --image dataset/1.jpg --checkpoint checkpoints/mask_rcnn_best.pth # 批量推理 + GT 对比 python predict.py --dataset dataset --show-gt """