Mask R-CNN instance segmentation training code
This commit is contained in:
@@ -0,0 +1,481 @@
|
||||
"""
|
||||
Mask R-CNN 实例分割训练脚本
|
||||
|
||||
用法:
|
||||
# 使用默认参数训练
|
||||
python train.py
|
||||
|
||||
# 自定义参数训练
|
||||
python train.py --dataset dataset --epochs 50 --batch-size 2 --lr 0.005
|
||||
|
||||
# 冻结骨干网络(小数据集推荐)
|
||||
python train.py --freeze-backbone --epochs 100
|
||||
|
||||
# 从检查点恢复训练
|
||||
python train.py --resume checkpoints/mask_rcnn_epoch_20.pth
|
||||
|
||||
功能:
|
||||
1. 自动扫描 dataset 目录中的 LabelMe 标注
|
||||
2. 支持 polygon / linestrip / rectangle / circle 标注类型
|
||||
3. 自动划分训练集/验证集
|
||||
4. 支持 COCO 预训练权重微调
|
||||
5. 支持冻结骨干网络(减少过拟合)
|
||||
6. 训练过程中保存检查点和最佳模型
|
||||
7. 记录训练日志(损失曲线)
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import json
|
||||
from typing import Dict, List
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from utils import get_label_map, get_image_json_pairs, split_dataset, visualize_instances
|
||||
from dataset import LabelMeDataset, get_train_transforms, get_val_transforms
|
||||
from model import get_model, freeze_backbone, get_optimizer, get_lr_scheduler
|
||||
|
||||
|
||||
def collate_fn(batch):
|
||||
"""
|
||||
自定义 collate 函数:Mask R-CNN 的每个样本可能有不同数量的实例,
|
||||
不能使用默认的 stack 方式,需要返回 list。
|
||||
"""
|
||||
return tuple(zip(*batch))
|
||||
|
||||
|
||||
def train_one_epoch(
|
||||
model: torch.nn.Module,
|
||||
optimizer: torch.optim.Optimizer,
|
||||
data_loader: DataLoader,
|
||||
device: torch.device,
|
||||
epoch: int,
|
||||
log_interval: int = 10,
|
||||
) -> Dict[str, float]:
|
||||
"""
|
||||
训练一个 epoch。
|
||||
|
||||
Args:
|
||||
model: 模型
|
||||
optimizer: 优化器
|
||||
data_loader: 训练数据加载器
|
||||
device: 计算设备 (cuda / cpu)
|
||||
epoch: 当前 epoch 编号
|
||||
log_interval: 日志打印间隔
|
||||
|
||||
Returns:
|
||||
各项损失的平均值
|
||||
"""
|
||||
model.train()
|
||||
total_loss = 0.0
|
||||
loss_components = {
|
||||
"loss_classifier": 0.0,
|
||||
"loss_box_reg": 0.0,
|
||||
"loss_mask": 0.0,
|
||||
"loss_objectness": 0.0,
|
||||
"loss_rpn_box_reg": 0.0,
|
||||
}
|
||||
num_batches = 0
|
||||
|
||||
try:
|
||||
from tqdm import tqdm
|
||||
pbar = tqdm(data_loader, desc=f"Epoch {epoch}")
|
||||
except ImportError:
|
||||
pbar = data_loader
|
||||
|
||||
for batch_idx, (images, targets) in enumerate(pbar):
|
||||
# 过滤无实例的样本
|
||||
valid_images = []
|
||||
valid_targets = []
|
||||
for img, tgt in zip(images, targets):
|
||||
if len(tgt["boxes"]) > 0:
|
||||
valid_images.append(img)
|
||||
valid_targets.append(tgt)
|
||||
|
||||
if len(valid_images) == 0:
|
||||
continue
|
||||
|
||||
images = [img.to(device) for img in valid_images]
|
||||
targets = [{k: v.to(device) for k, v in t.items()} for t in valid_targets]
|
||||
|
||||
# 前向传播(训练模式下返回 loss dict)
|
||||
loss_dict = model(images, targets)
|
||||
losses = sum(loss for loss in loss_dict.values())
|
||||
|
||||
# 反向传播
|
||||
optimizer.zero_grad()
|
||||
losses.backward()
|
||||
optimizer.step()
|
||||
|
||||
# 记录损失
|
||||
total_loss += losses.item()
|
||||
num_batches += 1
|
||||
for k, v in loss_dict.items():
|
||||
loss_components[k] += v.item()
|
||||
|
||||
# 更新进度条
|
||||
if hasattr(pbar, 'set_postfix'):
|
||||
pbar.set_postfix({"loss": f"{losses.item():.4f}"})
|
||||
|
||||
avg_loss = total_loss / max(num_batches, 1)
|
||||
for k in loss_components:
|
||||
loss_components[k] /= max(num_batches, 1)
|
||||
loss_components["total"] = avg_loss
|
||||
|
||||
return loss_components
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def evaluate(
|
||||
model: torch.nn.Module,
|
||||
data_loader: DataLoader,
|
||||
device: torch.device,
|
||||
) -> float:
|
||||
"""
|
||||
验证集评估:计算平均检测置信度作为近似指标。
|
||||
(完整的 COCO mAP 评估需要 pycocotools,这里使用简化指标避免额外依赖)
|
||||
|
||||
Args:
|
||||
model: 模型
|
||||
data_loader: 验证数据加载器
|
||||
device: 计算设备
|
||||
|
||||
Returns:
|
||||
平均置信度(0~1)
|
||||
"""
|
||||
model.eval()
|
||||
total_score = 0.0
|
||||
total_instances = 0
|
||||
|
||||
for images, targets in data_loader:
|
||||
images = [img.to(device) for img in images]
|
||||
outputs = model(images)
|
||||
|
||||
for output in outputs:
|
||||
scores = output["scores"]
|
||||
if len(scores) > 0:
|
||||
# 取置信度 > 0.5 的预测
|
||||
high_conf = scores[scores > 0.5]
|
||||
if len(high_conf) > 0:
|
||||
total_score += high_conf.mean().item()
|
||||
total_instances += 1
|
||||
|
||||
return total_score / max(total_instances, 1)
|
||||
|
||||
|
||||
def save_checkpoint(
|
||||
model: torch.nn.Module,
|
||||
optimizer: torch.optim.Optimizer,
|
||||
scheduler,
|
||||
epoch: int,
|
||||
label_map: Dict[str, int],
|
||||
loss_history: List[Dict],
|
||||
path: str,
|
||||
):
|
||||
"""保存训练检查点"""
|
||||
torch.save({
|
||||
"epoch": epoch,
|
||||
"model_state_dict": model.state_dict(),
|
||||
"optimizer_state_dict": optimizer.state_dict(),
|
||||
"scheduler_state_dict": scheduler.state_dict() if scheduler else None,
|
||||
"label_map": label_map,
|
||||
"loss_history": loss_history,
|
||||
}, path)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Mask R-CNN 实例分割训练")
|
||||
parser.add_argument("--dataset", type=str, default="dataset",
|
||||
help="数据集目录路径 (默认: dataset)")
|
||||
parser.add_argument("--output-dir", type=str, default="checkpoints",
|
||||
help="检查点输出目录 (默认: checkpoints)")
|
||||
parser.add_argument("--epochs", type=int, default=50,
|
||||
help="训练轮数 (默认: 50)")
|
||||
parser.add_argument("--batch-size", type=int, default=2,
|
||||
help="批大小 (默认: 2)")
|
||||
parser.add_argument("--lr", type=float, default=0.005,
|
||||
help="学习率 (默认: 0.005)")
|
||||
parser.add_argument("--momentum", type=float, default=0.9,
|
||||
help="SGD 动量 (默认: 0.9)")
|
||||
parser.add_argument("--weight-decay", type=float, default=0.0005,
|
||||
help="权重衰减 (默认: 0.0005)")
|
||||
parser.add_argument("--lr-step-size", type=int, default=10,
|
||||
help="学习率衰减间隔 (默认: 10)")
|
||||
parser.add_argument("--lr-gamma", type=float, default=0.1,
|
||||
help="学习率衰减系数 (默认: 0.1)")
|
||||
parser.add_argument("--val-ratio", type=float, default=0.2,
|
||||
help="验证集比例 (默认: 0.2)")
|
||||
parser.add_argument("--num-workers", type=int, default=0,
|
||||
help="数据加载线程数 (Windows 建议 0) (默认: 0)")
|
||||
parser.add_argument("--no-pretrained", action="store_true",
|
||||
help="不使用 COCO 预训练权重")
|
||||
parser.add_argument("--freeze-backbone", action="store_true",
|
||||
help="冻结骨干网络(小数据集推荐)")
|
||||
parser.add_argument("--resume", type=str, default=None,
|
||||
help="从检查点恢复训练的路径")
|
||||
parser.add_argument("--save-freq", type=int, default=5,
|
||||
help="每多少 epoch 保存一次检查点 (默认: 5)")
|
||||
parser.add_argument("--min-size", type=int, default=800,
|
||||
help="输入图像最小边尺寸 (默认: 800)")
|
||||
parser.add_argument("--max-size", type=int, default=1333,
|
||||
help="输入图像最大边尺寸 (默认: 1333)")
|
||||
parser.add_argument("--seed", type=int, default=42,
|
||||
help="随机种子 (默认: 42)")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# 设置随机种子
|
||||
torch.manual_seed(args.seed)
|
||||
np.random.seed(args.seed)
|
||||
|
||||
# 设置设备
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(f"使用设备: {device}")
|
||||
|
||||
# 构建标签映射
|
||||
dataset_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), args.dataset)
|
||||
label_map = get_label_map(dataset_path)
|
||||
print(f"标签映射: {label_map}")
|
||||
|
||||
num_classes = max(label_map.values()) + 1 # +1 for background
|
||||
print(f"类别数(含背景): {num_classes}")
|
||||
|
||||
# 获取数据配对并划分
|
||||
all_pairs = get_image_json_pairs(dataset_path)
|
||||
print(f"数据集总样本数: {len(all_pairs)}")
|
||||
|
||||
train_pairs, val_pairs = split_dataset(all_pairs, val_ratio=args.val_ratio, seed=args.seed)
|
||||
print(f"训练集: {len(train_pairs)} | 验证集: {len(val_pairs)}")
|
||||
|
||||
if len(train_pairs) == 0:
|
||||
print("错误: 训练集为空!请检查数据集目录。")
|
||||
sys.exit(1)
|
||||
|
||||
# 创建数据集和数据加载器
|
||||
train_dataset = LabelMeDataset(
|
||||
dataset_dir=dataset_path,
|
||||
label_map=label_map,
|
||||
pairs=train_pairs,
|
||||
transforms=get_train_transforms(),
|
||||
)
|
||||
val_dataset = LabelMeDataset(
|
||||
dataset_dir=dataset_path,
|
||||
label_map=label_map,
|
||||
pairs=val_pairs,
|
||||
transforms=get_val_transforms(),
|
||||
) if len(val_pairs) > 0 else None
|
||||
|
||||
train_loader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.batch_size,
|
||||
shuffle=True,
|
||||
num_workers=args.num_workers,
|
||||
collate_fn=collate_fn,
|
||||
drop_last=False,
|
||||
)
|
||||
val_loader = DataLoader(
|
||||
val_dataset,
|
||||
batch_size=1,
|
||||
shuffle=False,
|
||||
num_workers=args.num_workers,
|
||||
collate_fn=collate_fn,
|
||||
) if val_dataset else None
|
||||
|
||||
# 创建模型
|
||||
pretrained = not args.no_pretrained
|
||||
model = get_model(
|
||||
num_classes=num_classes,
|
||||
pretrained=pretrained,
|
||||
min_size=args.min_size,
|
||||
max_size=args.max_size,
|
||||
)
|
||||
print(f"预训练权重: {'COCO' if pretrained else '无'}")
|
||||
|
||||
if args.freeze_backbone:
|
||||
model = freeze_backbone(model, freeze=True)
|
||||
num_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
num_total = sum(p.numel() for p in model.parameters())
|
||||
print(f"骨干网络已冻结 | 可训练参数: {num_trainable:,} / {num_total:,}")
|
||||
else:
|
||||
num_total = sum(p.numel() for p in model.parameters())
|
||||
print(f"可训练参数: {num_total:,}")
|
||||
|
||||
model = model.to(device)
|
||||
|
||||
# 创建优化器和调度器
|
||||
optimizer = get_optimizer(model, lr=args.lr, momentum=args.momentum,
|
||||
weight_decay=args.weight_decay)
|
||||
scheduler = get_lr_scheduler(optimizer, step_size=args.lr_step_size, gamma=args.lr_gamma)
|
||||
|
||||
# 恢复训练
|
||||
start_epoch = 0
|
||||
loss_history = []
|
||||
best_score = 0.0
|
||||
|
||||
if args.resume and os.path.exists(args.resume):
|
||||
checkpoint = torch.load(args.resume, map_location=device, weights_only=False)
|
||||
model.load_state_dict(checkpoint["model_state_dict"])
|
||||
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
|
||||
if checkpoint.get("scheduler_state_dict") and scheduler:
|
||||
scheduler.load_state_dict(checkpoint["scheduler_state_dict"])
|
||||
start_epoch = checkpoint["epoch"] + 1
|
||||
loss_history = checkpoint.get("loss_history", [])
|
||||
best_score = max((h.get("val_score", 0) for h in loss_history), default=0.0)
|
||||
print(f"从 epoch {start_epoch} 恢复训练 (best_score={best_score:.4f})")
|
||||
|
||||
# 创建输出目录
|
||||
output_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), args.output_dir)
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# 保存标签映射
|
||||
with open(os.path.join(output_dir, "label_map.json"), "w", encoding="utf-8") as f:
|
||||
json.dump(label_map, f, ensure_ascii=False, indent=2)
|
||||
|
||||
# 训练循环
|
||||
print(f"\n{'='*60}")
|
||||
print(f"开始训练 | 总轮数: {args.epochs} | 批大小: {args.batch_size} | 学习率: {args.lr}")
|
||||
print(f"{'='*60}\n")
|
||||
|
||||
for epoch in range(start_epoch, args.epochs):
|
||||
epoch_start = time.time()
|
||||
|
||||
# 训练
|
||||
train_metrics = train_one_epoch(
|
||||
model, optimizer, train_loader, device, epoch + 1
|
||||
)
|
||||
|
||||
# 验证
|
||||
val_score = 0.0
|
||||
if val_loader is not None:
|
||||
val_score = evaluate(model, val_loader, device)
|
||||
|
||||
# 更新学习率
|
||||
scheduler.step()
|
||||
|
||||
epoch_time = time.time() - epoch_start
|
||||
current_lr = optimizer.param_groups[0]["lr"]
|
||||
|
||||
# 记录损失历史
|
||||
epoch_record = {
|
||||
"epoch": epoch + 1,
|
||||
"train_loss": train_metrics["total"],
|
||||
"loss_classifier": train_metrics["loss_classifier"],
|
||||
"loss_box_reg": train_metrics["loss_box_reg"],
|
||||
"loss_mask": train_metrics["loss_mask"],
|
||||
"loss_objectness": train_metrics["loss_objectness"],
|
||||
"loss_rpn_box_reg": train_metrics["loss_rpn_box_reg"],
|
||||
"val_score": val_score,
|
||||
"lr": current_lr,
|
||||
"time": round(epoch_time, 1),
|
||||
}
|
||||
loss_history.append(epoch_record)
|
||||
|
||||
# 打印日志
|
||||
print(
|
||||
f"Epoch {epoch+1}/{args.epochs} | "
|
||||
f"Loss: {train_metrics['total']:.4f} "
|
||||
f"(cls:{train_metrics['loss_classifier']:.4f} "
|
||||
f"box:{train_metrics['loss_box_reg']:.4f} "
|
||||
f"mask:{train_metrics['loss_mask']:.4f}) | "
|
||||
f"Val: {val_score:.4f} | "
|
||||
f"LR: {current_lr:.6f} | "
|
||||
f"Time: {epoch_time:.1f}s"
|
||||
)
|
||||
|
||||
# 保存最佳模型
|
||||
if val_score > best_score:
|
||||
best_score = val_score
|
||||
best_path = os.path.join(output_dir, "mask_rcnn_best.pth")
|
||||
save_checkpoint(model, optimizer, scheduler, epoch + 1, label_map,
|
||||
loss_history, best_path)
|
||||
print(f" -> 新最佳模型已保存 (score={best_score:.4f})")
|
||||
|
||||
# 定期保存检查点
|
||||
if (epoch + 1) % args.save_freq == 0 or epoch + 1 == args.epochs:
|
||||
ckpt_path = os.path.join(output_dir, f"mask_rcnn_epoch_{epoch+1}.pth")
|
||||
save_checkpoint(model, optimizer, scheduler, epoch + 1, label_map,
|
||||
loss_history, ckpt_path)
|
||||
print(f" -> 检查点已保存: {ckpt_path}")
|
||||
|
||||
# 保存最终模型
|
||||
final_path = os.path.join(output_dir, "mask_rcnn_final.pth")
|
||||
save_checkpoint(model, optimizer, scheduler, args.epochs, label_map,
|
||||
loss_history, final_path)
|
||||
print(f"\n训练完成!最终模型: {final_path}")
|
||||
|
||||
# 保存损失历史到 JSON
|
||||
with open(os.path.join(output_dir, "loss_history.json"), "w", encoding="utf-8") as f:
|
||||
json.dump(loss_history, f, ensure_ascii=False, indent=2)
|
||||
|
||||
# 绘制损失曲线
|
||||
try:
|
||||
_plot_loss_curve(loss_history, output_dir)
|
||||
except Exception as e:
|
||||
print(f"警告: 无法绘制损失曲线 ({e})")
|
||||
|
||||
print(f"最佳验证分数: {best_score:.4f}")
|
||||
|
||||
|
||||
def _plot_loss_curve(loss_history: List[Dict], output_dir: str):
|
||||
"""绘制并保存损失曲线图"""
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
epochs = [h["epoch"] for h in loss_history]
|
||||
total_loss = [h["train_loss"] for h in loss_history]
|
||||
cls_loss = [h["loss_classifier"] for h in loss_history]
|
||||
box_loss = [h["loss_box_reg"] for h in loss_history]
|
||||
mask_loss = [h["loss_mask"] for h in loss_history]
|
||||
val_scores = [h.get("val_score", 0) for h in loss_history]
|
||||
|
||||
fig, axes = plt.subplots(1, 2, figsize=(14, 5))
|
||||
|
||||
# 损失曲线
|
||||
ax1 = axes[0]
|
||||
ax1.plot(epochs, total_loss, "b-", label="Total Loss", linewidth=2)
|
||||
ax1.plot(epochs, cls_loss, "r--", label="Classifier Loss")
|
||||
ax1.plot(epochs, box_loss, "g--", label="Box Reg Loss")
|
||||
ax1.plot(epochs, mask_loss, "m--", label="Mask Loss")
|
||||
ax1.set_xlabel("Epoch")
|
||||
ax1.set_ylabel("Loss")
|
||||
ax1.set_title("Training Loss")
|
||||
ax1.legend()
|
||||
ax1.grid(True, alpha=0.3)
|
||||
|
||||
# 验证分数曲线
|
||||
ax2 = axes[1]
|
||||
ax2.plot(epochs, val_scores, "b-o", label="Val Score")
|
||||
ax2.set_xlabel("Epoch")
|
||||
ax2.set_ylabel("Score")
|
||||
ax2.set_title("Validation Score")
|
||||
ax2.legend()
|
||||
ax2.grid(True, alpha=0.3)
|
||||
|
||||
plt.tight_layout()
|
||||
curve_path = os.path.join(output_dir, "loss_curve.png")
|
||||
plt.savefig(curve_path, dpi=150, bbox_inches="tight")
|
||||
plt.close()
|
||||
print(f"损失曲线图已保存: {curve_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
"""
|
||||
# 训练(小数据集推荐冻结骨干网络)
|
||||
python train.py --freeze-backbone --epochs 100
|
||||
|
||||
# 推理
|
||||
python predict.py --image dataset/1.jpg --checkpoint checkpoints/mask_rcnn_best.pth
|
||||
|
||||
# 批量推理 + GT 对比
|
||||
python predict.py --dataset dataset --show-gt
|
||||
|
||||
"""
|
||||
Reference in New Issue
Block a user