Files

5.1 KiB
Raw Permalink Blame History

Mask R-CNN 实例分割训练框架

基于 PyTorch + torchvision 的 Mask R-CNN 目标分割训练代码,支持 LabelMe 标注格式。

功能特点

  • LabelMe 标注支持:自动解析 polygonlinestriprectanglecircle 四种标注类型
  • Linestrip 自动合并:将相同标签的多条线段智能合并为闭合多边形(适用于标注边界线而非直接画多边形的场景)
  • COCO 预训练微调:加载 COCO 预训练权重,仅替换分类头和 mask 预测头,加速收敛
  • 数据增强:随机水平翻转、亮度调整、对比度调整(同时变换 mask 和 boxes)
  • 训练管理:检查点保存/恢复、最佳模型追踪、损失曲线可视化
  • 推理可视化:支持单图/批量推理,叠加显示 mask 和包围框

环境要求

  • Python 3.11
  • Windows / Linux / macOS

安装依赖

# 激活虚拟环境(如已有)
# .\env\Scripts\activate

# 安装 PyTorch (CPU 版本)
pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu

# 安装其他依赖
pip install opencv-python tqdm

# 或一次性安装
pip install -r requirements.txt

项目结构

segmentation/
├── dataset/              # 数据集目录
│   ├── 1.jpg             # 图像
│   ├── 1.json           # LabelMe 标注
│   ├── 2.jpg
│   ├── 2.json
│   └── ...
├── utils.py              # 工具函数(标注解析、mask 转换、可视化)
├── dataset.py           # Dataset 类和数据增强
├── model.py             # Mask R-CNN 模型构建
├── train.py             # 训练脚本
├── predict.py           # 推理脚本
├── requirements.txt     # 依赖列表
└── README.md            # 本文件

使用方法

1. 训练模型

# 默认参数训练(使用 COCO 预训练权重)
python train.py

# 自定义参数
python train.py --epochs 100 --batch-size 2 --lr 0.005

# 冻结骨干网络(小数据集推荐,减少过拟合)
python train.py --freeze-backbone --epochs 100

# 调整图像输入尺寸
python train.py --min-size 512 --max-size 800

# 从检查点恢复训练
python train.py --resume checkpoints/mask_rcnn_epoch_20.pth

主要参数说明:

参数 默认值 说明
--dataset dataset 数据集目录路径
--epochs 50 训练轮数
--batch-size 2 批大小
--lr 0.005 学习率
--val-ratio 0.2 验证集比例
--freeze-backbone False 冻结骨干网络
--no-pretrained False 不使用预训练权重
--resume None 恢复训练的检查点路径

2. 推理与可视化

# 对单张图片推理
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 --threshold 0.7

# 同时显示 ground truth 对比
python predict.py --dataset dataset --show-gt

数据集格式

使用 LabelMe 工具进行标注,数据目录结构:

dataset/
├── image1.jpg
├── image1.json
├── image2.jpg
├── image2.json
└── ...

每个 JSON 文件的 shapes 列表中每个 shape 包含:

  • label: 类别名称(如 "crosswalk"
  • points: [[x, y], ...] 坐标点列表
  • shape_type: 标注类型

标注类型说明

类型 说明 处理方式
polygon 多边形 直接转换为 mask
linestrip 线段 相同标签的线段自动合并为多边形
rectangle 矩形 转换为四点多边形
circle 圆形 近似为多边形

Linestrip 合并:当标注对象边界使用多条线段(而非直接画多边形)时, 代码会自动将相同标签的线段按最近端点连接,合并为闭合多边形。

训练输出

训练完成后,checkpoints/ 目录包含:

文件 说明
mask_rcnn_best.pth 验证分数最高的模型
mask_rcnn_final.pth 最后一轮的模型
mask_rcnn_epoch_N.pth 每 N 轮的检查点
label_map.json 标签映射表
loss_history.json 损失历史记录
loss_curve.png 损失曲线图

小数据集建议

当前数据集仅有 4 张图片,建议:

  1. 使用 --freeze-backbone 冻结骨干网络,大幅减少可训练参数
  2. 增加 epoch 数(如 --epochs 200),因为样本少需要更多迭代
  3. 增加数据:4 张图片不足以训练出泛化能力强的模型,建议扩充到至少 100+ 张
  4. 使用 polygon 标注:直接画多边形比 linestrip 更精确

技术细节

  • 模型Mask R-CNN + ResNet50-FPN 骨干网络(torchvision 实现)
  • 优化器SGD + momentum(0.9) + weight_decay(0.0005)
  • 学习率调度StepLR,每 10 个 epoch 衰减为 0.1 倍
  • 数据增强:水平翻转 + 亮度/对比度抖动(同时变换 mask 和 boxes)