482 lines
16 KiB
Python
482 lines
16 KiB
Python
"""
|
||
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
|
||
|
||
"""
|