训练何时提前停止,应看独立验证数据上的目标指标,而不是只看训练损失还在不在下降。对“越小越好”的验证损失,可以设定改善阈值 min_delta 和等待次数 patience,连续若干次验证没有足够改善就停止,并恢复保存的最佳模型参数。
下面用普通 PyTorch 循环实现这一规则,不依赖 Trainer。Transformers 早停回调说明中,patience 也是以评估调用次数计数;本文是自编实现,保存“实际最低损失”与判断“有意义改善”采用两个独立变量,不能把它当作该回调的逐行复刻。

先约定判断口径
- 每完成一轮参数更新,验证一次;因此本例一次检查正好对应一轮。每隔多轮才验证时,patience 仍按验证次数数。
- best_seen 保存所有已见验证损失的最低值,用于选择最后恢复的参数。
- progress_reference 记录最近一次超过 min_delta 的有效进步,用来控制等待计数。
- 验证损失不是有限数时直接报错,不能当作普通“未进步”继续训练。
小于阈值的进步仍可能产生真正最好的模型,所以不应只在重置 patience 时保存参数。最佳状态需要深拷贝;只把 model.state_dict() 的引用放到变量里,后续更新可能让你失去当时的参数。PyTorch 保存教程也明确提醒保存最佳模型时使用深拷贝或序列化。
跑一段会提前停止的完整代码
以下完整脚本在 Windows、Python 3.11.15、PyTorch 2.8.0+cpu、NumPy 2.3.3 环境实际运行通过,只验证这个 CPU 玩具任务。它不下载模型权重,也不代表真实模型的精度、速度或显存表现。在已启用的独立 Python 环境中可安装相同依赖:
python -m pip install torch==2.8.0 --index-url https://download.pytorch.org/whl/cpu
python -m pip install numpy==2.3.3
本例故意把训练标签设置为 y=3x+0.5、验证标签设置为 y=2x+0.5,制造后续验证损失恶化,检查早停和恢复逻辑。这是人为构造的压力样例,不是证明模型泛化表现的实验;正式训练需使用任务定义一致、没有泄漏的真实训练集和验证集。
保存为 early_stop_demo.py,运行 python early_stop_demo.py。patience=3、min_delta=0.001 只是本例的演示值,不能直接当成所有任务的推荐参数。
import copy
import math
import torch
torch.set_num_threads(1)
torch.manual_seed(7)
x_train = torch.linspace(-1, 1, 32).reshape(-1, 1)
y_train = 3 * x_train + 0.5 # Artificial label shift to exercise early stopping.
x_valid = torch.linspace(-0.9, 0.9, 16).reshape(-1, 1)
y_valid = 2 * x_valid + 0.5
model = torch.nn.Linear(1, 1)
optimizer = torch.optim.SGD(model.parameters(), lr=0.2)
criterion = torch.nn.MSELoss()
patience = 3
min_delta = 0.001 # Demonstration threshold; not a tuned recommendation.
best_seen = math.inf
progress_reference = math.inf
bad_checks = 0
best_state = None
best_epoch = None
for epoch in range(1, 101):
model.train()
optimizer.zero_grad(set_to_none=True)
loss = criterion(model(x_train), y_train)
loss.backward()
optimizer.step()
model.eval()
with torch.no_grad():
valid_loss = criterion(model(x_valid), y_valid).item()
if not math.isfinite(valid_loss):
raise RuntimeError("non-finite validation loss")
# Save the actual best result, even when the improvement is small.
if valid_loss < best_seen:
best_seen = valid_loss
best_state = copy.deepcopy(model.state_dict())
best_epoch = epoch
# Patience counts validation calls without a meaningful improvement.
if valid_loss < progress_reference - min_delta:
progress_reference = valid_loss
bad_checks = 0
else:
bad_checks += 1
if bad_checks >= patience:
print("stop_epoch:", epoch)
break
last_loss = valid_loss
model.load_state_dict(best_state)
model.eval()
with torch.no_grad():
restored_loss = criterion(model(x_valid), y_valid).item()
print("best_epoch:", best_epoch)
print("best_validation_loss:", round(best_seen, 8))
print("last_loss_worse:", last_loss > best_seen)
print("restored_loss_matches:", abs(restored_loss - best_seen) < 1e-12)
assert abs(restored_loss - best_seen) < 1e-12
assert last_loss > best_seen
确认恢复的是最佳轮次
本次 CPU 运行输出:
stop_epoch: 10
best_epoch: 7
best_validation_loss: 0.00095615
last_loss_worse: True
restored_loss_matches: True
第 10 轮触发停止,第 7 轮是实际验证损失最低的轮次。last_loss_worse 为 True 说明停下来的最后参数已经更差;restored_loss_matches 为 True 表示重新加载最佳状态后,验证损失与记录的最佳值一致。读者应核对停止轮次、最佳轮次和恢复后的重新评估结果,不能把“break 执行了”当成全部验收。
实际项目怎样接入
- 保留自己原来的训练循环,用独立验证集计算每次验证的平均指标;各轮采用相同样本和相同计算口径。
- 验证时调用 model.eval(),用 torch.no_grad() 关闭梯度;下一轮训练前重新调用 model.train()。
- 按指标方向改比较条件。损失通常越小越好,准确率等越大越好,不能照抄相同不等式。
- 训练结束后加载 best_state,再在独立测试集做最终评估。验证集用来调训练策略,不能同时当成最终无偏评估依据。
完整模型较大时,可把最佳 state_dict 保存到专用磁盘路径,降低复制一份参数对内存的压力。PyTorch 迁移学习教程展示了按验证指标保存并在结束后加载最佳权重的流程。
常见问题:停得太早,还是迟迟不停?
先检查验证频率、指标方向、样本是否每次一致,再看 min_delta 是否相对指标量级过大,以及噪声是否让 patience 不断重置。本文仅恢复用于推理或后续评估的最佳模型参数;若要从最佳轮次继续训练,还要一起保存那个轮次的优化器与其他训练状态,不能配上最后一轮优化器直接称为完整恢复。
Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30312.html