Mask R-CNN instance segmentation training code

This commit is contained in:
flk
2026-08-10 01:20:28 +08:00
commit 61f6ea918e
16 changed files with 2026 additions and 0 deletions
+481
View File
@@ -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
"""