162 lines
5.1 KiB
Markdown
162 lines
5.1 KiB
Markdown
|
|
# 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
|
|||
|
|
|
|||
|
|
## 安装依赖
|
|||
|
|
|
|||
|
|
```bash
|
|||
|
|
# 激活虚拟环境(如已有)
|
|||
|
|
# .\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. 训练模型
|
|||
|
|
|
|||
|
|
```bash
|
|||
|
|
# 默认参数训练(使用 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. 推理与可视化
|
|||
|
|
|
|||
|
|
```bash
|
|||
|
|
# 对单张图片推理
|
|||
|
|
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](https://github.com/wkentaro/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)
|