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
|
|||
|
|
|
|||
|
|
"""
|