You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.22 01:06:20