6.早停机制
第 6 章 — 早停机制(Early Stopping)
第 5 章我们学会了把"训练现场"完整存进 checkpoint,还能用 best.pth 留住验证集表现最好的那个模型。但训练本身仍然"闷头跑满 --epochs 指定的全部轮数":如果模型第 20 轮就到顶了,剩下 30 轮既浪费时间,还可能让模型在训练集上继续"死记硬背"。本章加一个早停机制(Early Stopping):验证集指标连续 N 轮不再提升,训练就自动提前收工—— 由机器来喊停,而不是人肉盯日志按 Ctrl+C 。
一、本章要解决的问题
- 第 5 章结束后的现状:
--epochs 100时若 40 轮就到顶,剩下 60 轮照样会跑完——纯浪费时间,还平白增加过拟合风险。 - 训练轮数很难提前猜准:设少了欠拟合(没学够),设多了浪费时间 + 过拟合(学过头)。
- 手动盯着终端日志、看到曲线走平就 Ctrl+C:不可靠、不可复现,而且人总有走神的时候。
- 判断标准:训练能在"模型不再变好"时自动终止,并且终止后仍能用
best.pth取出打得最好的那个模型。
二、核心概念速览
下面 7 个概念是本章代码的全部"生词"。先花 5 分钟读完,再看代码会轻松很多。
1. 过拟合(死记硬背答案的学生)
把训练比作学生备考:训练集是"练习册",模型就是反复刷题的学生。真学会的学生,换一份从没见过的试卷(验证集)也能做对;只会死记硬背的学生,练习册背得滚瓜烂熟,一换题就露馅。过拟合就是这种"背答案"状态:训练集准确率一路狂涨,验证集准确率却不再涨、甚至开始回落——因为模型开始把训练数据里的噪声当成了规律。早停要防的正是它。
2. 验证集在模型选择中的作用(考试模拟卷)
训练题是练习册(训练集),最后的考卷是测试集,而测试卷只能用一次(第 2 章就讲过)。那平时怎么判断"这学生现在行不行"?用模拟卷——验证集。它和训练集完全不重叠,模型从没抄过它的答案,所以它的分数比训练集分数更能反映"真本事"。早停"拿哪个分数做决策"?拿验证集分数,绝不能用训练集分数——用训练集分数做决策等于自己考自己。
3. patience(给几次机会)
训练曲线不是一路直上的,中途小波动很正常:这一轮没涨,下一轮可能就突破了。所以"一轮不涨就停"会误杀还在爬坡的模型。patience 的意思就是给几次机会:允许连续 N 轮不进步,如果连续 N 轮都没刷新纪录,才判定"真的到头了"。N 越大越保守(不容易误停,但可能多跑很多轮),N 越小越激进(省时间,但容易误杀),一般取总 epoch 数的 10%~20% 起步。
4. best_score 与 counter(历史最佳成绩与"摆烂计数器")
早停的全部状态就这两样:best_score 记"见过的最好 val_acc"(班级历史最高分),counter 数"连续几轮没超过这个最高分"(摆烂天数)。每轮验证完做一次判定:破了纪录 → best_score 更新、counter 归零;没破纪录 → counter +1。当 counter 顶到 patience,判定"真的不行了",触发早停。一个记巅峰、一个数天数,两者配合就是完整的早停逻辑。
5. delta 容差(宽容"毫厘之差")
best_score 是 0.8000,下一轮 0.7999,算退步吗?严格说算,但这种小数点后第四位的差异纯粹是噪声,不该因此重置计数。delta 就是一个 容差阈值 :只有当 val_acc >= best_score + delta 才算真正刷新纪录,否则一律按"未提升"计数。delta=0 是最严格版本(必须严格更高),delta=0.001 意味着"涨 0.1 个百分点才算进步"。
6. 早停与 best.pth 判定逻辑的关系(同一个裁判,两个用途)
第 5 章的 if val_acc > best_acc 才保存 best.pth,本质是"用验证集分数给每轮打分,留下最高分的存档";本章 EarlyStopping 里的 best_score 判定同样是在看验证集分数有没有创新高—— 同一个裁判,两个用途 :best.pth 回答"什么时候该把当前模型存下来",早停回答"什么时候该整体收工"。这也是为什么早停触发后做测试评估应该加载 best.pth 而不是 last.pth:停下来的那个模型,不一定是最好的那个模型。
7. 早停与 cosine 调度搭配的注意点(提前下车会打断余弦曲线)
CosineAnnealingLR 的整个设计前提是"跑满 T_max 个 epoch":学习率沿一条余弦曲线先快降再慢降,最后几个 epoch 用极小的学习率做精细打磨。如果早停在曲线中段就提前收工,后半段"低学习率精调"全被浪费,还可能停在模型尚未收敛充分的位置。所以早停与 cosine 搭配时建议把 patience 调大(给余弦曲线留足走完的时间),而 StepLR 这种"按固定步长衰减"的调度跟早停搭配更自然。
三、解决思路
- 封装一个
EarlyStopping类:把"记最好成绩、数没提升轮数"这两件事收敛成一个对象,每轮验证后调用一次early_stopping(val_acc)即可,主循环几乎不加负担。 - 判定规则:
score < best_score + delta视为未提升 →counter += 1;否则刷新best_score、counter归零。当counter >= patience时置early_stop = True。 - 与第 5 章的存档机制无缝衔接:每轮照样存
last.pth、刷新纪录照样存best.pth;早停只是给主循环加一个break,存档逻辑一行没动。 - 参数化:新增
--patience命令行参数,方便不同任务调整"给几次机会";--epochs默认值调大到 30,为早停留足"可提前终止"的空间——如果 epochs 本来就很小,早停多半不会触发。
trade-off:早停是用验证集做"模型选择"的自动化,代价是验证集被反复用来做决策,统计上会产生一点"选择偏置"(最终 test acc 通常比 val acc 略低),对 CIFAR-10 这种规模完全可接受。
patience太小容易在爬坡中被误杀,太大则失去省时意义,一般取总 epoch 数的 10%~20% 起步。
四、代码变更
相对第 5 章的改动(tqdm 自第 3 章引入、checkpoint 自第 5 章引入,本章均沿用不动):
- parser.add_argument("--epochs", type=int, default=10, help="训练轮数")
+ parser.add_argument("--epochs", type=int, default=30, help="训练轮数上限")
+ # 默认 10 -> 30:早停的前提是"留足训练空间",让它有机会在中途自动喊停
parser.add_argument("--batch-size", type=int, default=64, help="每个 batch 的样本数")
...
parser.add_argument("--resume", type=str, default=None,
help="从指定 checkpoint 恢复训练")
+ parser.add_argument("--patience", type=int, default=7,
+ help="早停:验证集连续多少轮无提升则停止")
...
+ # ----------------------------------------------------------------------------
+ # 新增:EarlyStopping 早停工具类——核心就是 best_score 与 counter 两个状态
+ # ----------------------------------------------------------------------------
+ class EarlyStopping:
+ """验证集指标连续 patience 轮无提升时,置 early_stop=True 通知主循环停止。"""
+ def __init__(self, patience=7, delta=0.0, verbose=False):
+ self.patience = patience
+ self.delta = delta
+ self.verbose = verbose
+ self.best_score = None # 历史最优 val_acc
+ self.counter = 0 # 连续未提升的轮数
+ self.early_stop = False
+
+ def __call__(self, val_acc):
+ score = val_acc
+ if self.best_score is None:
+ self.best_score = score # 第一轮:直接记为历史最优
+ elif score < self.best_score + self.delta: # 含 delta 容差
+ self.counter += 1 # 没创新高:摆烂天数 +1
+ if self.verbose:
+ print(f"[EarlyStopping] 连续 {self.counter}/{self.patience} 轮未提升")
+ if self.counter >= self.patience:
+ self.early_stop = True # 连续太多轮:通知主循环收工
+ else:
+ self.best_score = score # 创新高:更新纪录,计数归零
+ self.counter = 0
+
+ early_stopping = EarlyStopping(patience=args.patience, verbose=True)
+
for epoch in range(start_epoch, args.epochs + 1):
...
if val_acc > best_acc:
best_acc = val_acc
state["best_acc"] = best_acc
save_checkpoint(state, os.path.join(args.ckpt_dir, "best.pth"))
+
+ # 早停判断:无提升达到 patience 轮则提前结束
+ early_stopping(val_acc)
+ if early_stopping.early_stop:
+ print(f"[EarlyStopping] 连续 {args.patience} 轮无提升,在 epoch {epoch} 停止")
+ break除此之外,train_one_epoch / validate 的 tqdm 进度条(第 3 章引入)与 last.pth / best.pth 双文件存档逻辑(第 5 章引入)均原样保留,本章不碰它们。
五、完整代码
创建 train.py(完整版):
import os
import argparse
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
from collections import defaultdict
from tqdm import tqdm # 进度条:第 3 章引入,本章沿用
def parse_args():
parser = argparse.ArgumentParser(description="CIFAR-10 图像分类训练")
parser.add_argument("--epochs", type=int, default=30, help="训练轮数上限")
parser.add_argument("--batch-size", type=int, default=64, help="每个 batch 的样本数")
parser.add_argument("--lr", type=float, default=1e-3, help="初始学习率")
parser.add_argument("--lr-scheduler", type=str, default="step",
choices=["step", "cosine"], help="学习率调度策略")
parser.add_argument("--lr-step-size", type=int, default=15, help="StepLR 每多少轮衰减一次")
parser.add_argument("--lr-gamma", type=float, default=0.1, help="StepLR 衰减系数")
parser.add_argument("--data-dir", type=str, default="./data", help="数据集存放目录")
parser.add_argument("--num-workers", type=int, default=2, help="DataLoader 数据加载进程数")
parser.add_argument("--device", type=str, default="auto", help="运行设备:auto/cuda/cpu")
parser.add_argument("--ckpt-dir", type=str, default="./checkpoints", help="checkpoint 保存目录")
parser.add_argument("--resume", type=str, default=None, help="从指定 checkpoint 恢复训练")
# 本章新增参数:连续多少轮验证集无提升就提前终止训练
parser.add_argument("--patience", type=int, default=7, help="早停:验证集连续多少轮无提升则停止")
return parser.parse_args()
args = parse_args()
# 设备选择:与第 5 章完全一致,有 GPU 用 GPU,否则用 CPU
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
print(f"Using device: {device}")
# ----------------------------------------------------------------------------
# 模型:SimpleCNN 与前几章完全一致
# 3 个卷积块(提特征)+ 全局平均池化 + 线性分类头(分类)
# ----------------------------------------------------------------------------
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2), # 32x32 -> 16x16
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2), # 16x16 -> 8x8
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.MaxPool2d(2), # 8x8 -> 4x4
)
self.classifier = nn.Sequential(
nn.AdaptiveAvgPool2d(1), # 任意输入尺寸 -> 1x1,免去手算展平维数
nn.Flatten(), # (B,128,1,1) -> (B,128)
nn.Linear(128, num_classes),
)
def forward(self, x):
return self.classifier(self.features(x))
# ----------------------------------------------------------------------------
# 数据:CIFAR-10,训练集切成 45000 训练 + 5000 验证,测试集 10000
# ----------------------------------------------------------------------------
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
])
train_dataset = datasets.CIFAR10(root=args.data_dir, train=True, download=True, transform=transform)
test_dataset = datasets.CIFAR10(root=args.data_dir, train=False, download=True, transform=transform)
train_dataset, val_dataset = random_split(
train_dataset, [45000, len(train_dataset) - 45000]
)
train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
shuffle=True, num_workers=args.num_workers)
val_loader = DataLoader(val_dataset, batch_size=args.batch_size,
shuffle=False, num_workers=args.num_workers)
test_loader = DataLoader(test_dataset, batch_size=args.batch_size,
shuffle=False, num_workers=args.num_workers)
# ----------------------------------------------------------------------------
# 损失函数 + 优化器 + 学习率调度器(step / cosine 二选一,与第 4 章一致)
# ----------------------------------------------------------------------------
model = SimpleCNN().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=args.lr)
if args.lr_scheduler == "step":
scheduler = optim.lr_scheduler.StepLR(
optimizer, step_size=args.lr_step_size, gamma=args.lr_gamma)
elif args.lr_scheduler == "cosine":
scheduler = optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=args.epochs)
else:
scheduler = None
# ----------------------------------------------------------------------------
# 训练 / 验证 / 存档:与第 5 章一致,本章不修改任何逻辑
# ----------------------------------------------------------------------------
def train_one_epoch(model, loader, criterion, optimizer, device):
model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(loader, desc="Train", leave=False):
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def validate(model, loader, criterion, device):
model.eval()
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad():
for images, labels in tqdm(loader, desc="Val", leave=False):
images, labels = images.to(device), labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def save_checkpoint(state, path):
os.makedirs(os.path.dirname(path), exist_ok=True)
torch.save(state, path)
print(f"[Checkpoint] 已保存到 {path}")
# ----------------------------------------------------------------------------
# 本章新内容:EarlyStopping 早停工具类
# 核心就两个状态:best_score(历史最优 val_acc)+ counter(连续未提升轮数)
# ----------------------------------------------------------------------------
class EarlyStopping:
"""验证集指标连续 patience 轮无提升时,置 early_stop=True 通知主循环停止。"""
def __init__(self, patience=7, delta=0.0, verbose=False):
self.patience = patience
self.delta = delta
self.verbose = verbose
self.best_score = None # 历史最优 val_acc
self.counter = 0 # 连续未提升的轮数
self.early_stop = False
def __call__(self, val_acc):
score = val_acc
if self.best_score is None:
self.best_score = score # 第一轮:直接记为历史最优
elif score < self.best_score + self.delta: # 含 delta 容差
self.counter += 1 # 没创新高:摆烂天数 +1
if self.verbose:
print(f"[EarlyStopping] 连续 {self.counter}/{self.patience} 轮未提升")
if self.counter >= self.patience:
self.early_stop = True # 连续太多轮无提升:通知主循环收工
else:
self.best_score = score # 创新高:更新纪录,计数归零
self.counter = 0
# 恢复逻辑:--resume 指定存档时无缝续跑(与第 5 章一致)
metrics = defaultdict(list)
start_epoch, best_acc = 1, 0.0
if args.resume:
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model_state"])
optimizer.load_state_dict(ckpt["optimizer_state"])
if scheduler is not None and ckpt.get("scheduler_state") is not None:
scheduler.load_state_dict(ckpt["scheduler_state"])
start_epoch = ckpt["epoch"] + 1
best_acc = ckpt.get("best_acc", 0.0)
metrics = defaultdict(list, ckpt.get("metrics", {}))
print(f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={best_acc:.4f}")
# 实例化早停器:patience 从命令行读入,verbose=True 打印每次"未提升"提示
early_stopping = EarlyStopping(patience=args.patience, verbose=True)
print("=" * 60)
print("训练配置:")
for k, v in vars(args).items():
print(f" {k:12s} = {v}")
print("=" * 60)
for epoch in range(start_epoch, args.epochs + 1):
train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
val_loss, val_acc = validate(model, val_loader, criterion, device)
if scheduler is not None:
scheduler.step()
current_lr = optimizer.param_groups[0]["lr"]
metrics["train_loss"].append(train_loss)
metrics["train_acc"].append(train_acc)
metrics["val_loss"].append(val_loss)
metrics["val_acc"].append(val_acc)
print(f"Epoch {epoch:02d}/{args.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
# --- 存档:last.pth 每轮都存,best.pth 只在刷新验证集最优时存(第 5 章逻辑)---
state = {
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"scheduler_state": scheduler.state_dict() if scheduler else None,
"epoch": epoch,
"best_acc": best_acc,
"metrics": dict(metrics),
}
save_checkpoint(state, os.path.join(args.ckpt_dir, "last.pth"))
if val_acc > best_acc:
best_acc = val_acc
state["best_acc"] = best_acc
save_checkpoint(state, os.path.join(args.ckpt_dir, "best.pth"))
# 早停判断:无提升达到 patience 轮则提前结束
early_stopping(val_acc)
if early_stopping.early_stop:
print(f"[EarlyStopping] 连续 {args.patience} 轮无提升,在 epoch {epoch} 停止")
break
# 训练结束后在测试集上评估一次
# 注意:早停触发时,正式评估应当加载 best.pth(本章小结会再次强调)
test_loss, test_acc = validate(model, test_loader, criterion, device)
print(f"\n[Test] loss {test_loss:.4f} acc {test_acc:.4f}")六、本章小结
- 学到了什么
- 早停 = 用验证集做"模型选择"的自动化,核心只有两个状态:
best_score(历史最优 val_acc)与counter(连续未提升轮数)。 patience是"给几次机会",delta是"容差阈值",两者共同防止因微小波动误停。- 早停与第 5 章
best.pth判定同源(都用 val_acc 做裁决):best.pth负责"何时存下当前模型",早停负责"何时整体收工"。
- 早停 = 用验证集做"模型选择"的自动化,核心只有两个状态:
- 常见坑
counter数的是"连续未提升"的轮数,不是累计——期间任何一次刷新纪录都会归零,别把它当成"总停滞次数"。- 早停触发后做测试评估,务必加载
best.pth而不是last.pth:停下来的那个模型不一定是最好的那个。 - 与 cosine 调度搭配时,早停可能打断余弦曲线后半段的低学习率精调;要么调大
patience,要么干脆用 StepLR。 - 验证集被反复用于决策会产生轻微"选择偏置",最终 test acc 略低于 val acc 属正常现象,不用慌。
- 下一章预告:目前日志全靠
print打在终端里,一闪而过——想看某个 epoch 的指标变化只能翻屏幕,想画 loss 曲线还得手动把数据抄进 Excel。第 7 章引入logging模块 + TensorBoard,让训练过程可记录、可回溯、可可视化。
七、动手练习
- 观察停止时机:
--epochs 100 --patience 3跑一次,记下"早停触发的 epoch"与best.pth里记录的 epoch,两者差多少?思考为什么best.pth的 epoch 一定 ≤ 早停触发的 epoch。 - 调 delta:把
EarlyStopping的delta设为 0.001,对比停止时机变化;想一想什么场景适合非零delta(提示:验证集随机性带来的微小波动)。 - 早停 + 最优模型:早停触发后,分别加载
best.pth与last.pth在测试集上评估(复用第 5 章的评估脚本),对比两者的 test acc,亲身体会"停下来时的模型 ≠ 打得最好的模型"。 - 对比调度器:相同
patience下分别用--lr-scheduler step与cosine跑一轮,观察早停触发的 epoch 差异,体会"余弦曲线被提前打断"的实际影响。
