Files
segmentation/README.md
T

162 lines
5.1 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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