PyTorch Lightning:损失意外飙升时如何自动重载最新检查点?
解决方案:自动恢复损失飙升后的训练流程
核心思路是在训练循环中实时监控损失值,当检测到异常飙升时,触发加载检查点、重置优化器的逻辑。以下以PyTorch为例给出落地方案(TensorFlow可参照相同思路实现):
1. 损失异常检测逻辑
先明确损失飙升的判定规则,比如:
- 当前损失超过正常训练时损失的3倍以上
- 当前损失较上一步骤突增5倍及以上
示例代码:
# 初始化上一步损失 prev_loss = float('inf') # 设定异常判定阈值 LOSS_SPIKE_RATIO = 5.0 # 当前损失是上一步的5倍则判定为飙升 def is_loss_spiked(current_loss): global prev_loss if prev_loss == float('inf'): prev_loss = current_loss return False spike_ratio = current_loss / prev_loss prev_loss = current_loss return spike_ratio > LOSS_SPIKE_RATIO
2. 检查点保存机制
训练过程中定期保存模型权重、优化器状态及当前训练进度,确保能恢复到最近的正常状态:
import torch def save_checkpoint(model, optimizer, epoch, step, save_path='./latest_checkpoint.pth'): checkpoint = { 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'epoch': epoch, 'step': step, 'last_loss': prev_loss # 保存上一步的正常损失值 } torch.save(checkpoint, save_path)
建议每N个epoch或固定步数自动保存一次,也可在损失下降时额外保存最优检查点。
3. 异常触发后的恢复流程
检测到损失飙升时,加载最新检查点,重置优化器并从保存的进度继续训练:
def load_checkpoint(model, optimizer, load_path='./latest_checkpoint.pth'): checkpoint = torch.load(load_path) model.load_state_dict(checkpoint['model_state_dict']) # 重置优化器:加载原状态后可选择重新初始化,或直接复用加载的状态 optimizer.load_state_dict(checkpoint['optimizer_state_dict']) # 若使用学习率调度器,需同步恢复其状态 # scheduler.load_state_dict(checkpoint['scheduler_state_dict']) epoch = checkpoint['epoch'] step = checkpoint['step'] global prev_loss prev_loss = checkpoint['last_loss'] return epoch, step
4. 整合到训练循环
将上述逻辑嵌入训练主循环:
# 初始化模型、优化器 model = YourModel() optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) # 训练循环初始化 start_epoch = 0 start_step = 0 total_epochs = 100 for epoch in range(start_epoch, total_epochs): model.train() for step, (data, target) in enumerate(dataloader, start=start_step): optimizer.zero_grad() output = model(data) loss = your_loss_fn(output, target) loss.backward() optimizer.step() # 检查损失是否异常飙升 if is_loss_spiked(loss.item()): print(f"Loss spiked at epoch {epoch}, step {step}! Restoring from checkpoint...") # 恢复检查点并更新训练起始位置 start_epoch, start_step = load_checkpoint(model, optimizer) # 跳出当前step循环,从恢复后的位置继续 break else: # 当前epoch无异常,保存最新检查点 save_checkpoint(model, optimizer, epoch, step) start_step = 0 # 下一个epoch从step 0开始
额外优化建议
- 动态调整学习率:恢复训练后临时降低学习率(比如乘以0.5),减少再次触发异常的概率。
- 梯度裁剪:在
loss.backward()后添加torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0),防止梯度爆炸引发损失飙升。 - 多版本检查点:保存最近5个历史检查点,避免最新检查点也存在异常的情况。
内容的提问来源于stack exchange,提问作者Rylan Schaeffer
相关产品推荐
相关产品推荐

