5.1 KiB
5.1 KiB
Mask R-CNN 实例分割训练框架
基于 PyTorch + torchvision 的 Mask R-CNN 目标分割训练代码,支持 LabelMe 标注格式。
功能特点
- LabelMe 标注支持:自动解析
polygon、linestrip、rectangle、circle四种标注类型 - 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 张图片,建议:
- 使用
--freeze-backbone冻结骨干网络,大幅减少可训练参数 - 增加 epoch 数(如
--epochs 200),因为样本少需要更多迭代 - 增加数据:4 张图片不足以训练出泛化能力强的模型,建议扩充到至少 100+ 张
- 使用 polygon 标注:直接画多边形比 linestrip 更精确
技术细节
- 模型:Mask R-CNN + ResNet50-FPN 骨干网络(torchvision 实现)
- 优化器:SGD + momentum(0.9) + weight_decay(0.0005)
- 学习率调度:StepLR,每 10 个 epoch 衰减为 0.1 倍
- 数据增强:水平翻转 + 亮度/对比度抖动(同时变换 mask 和 boxes)