Files

162 lines
5.1 KiB
Markdown
Raw Permalink Normal View History

# 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