多卡场景中加载半月前checkpoint后loss突然变大的调试思路

文章导读
模型训练时遇到「加载旧 checkpoint 后 loss 突然变大」,先别急着判断权重损坏。“半月前”这个时间差本身就是一条重要线索:这期间代码、数据、依赖甚至训练配置都可能改过。建议把排查拆成两个子问题:checkpoint 能否恢复到保存时的状态;当前环境是否还能按原来的方式组织样本和计算 loss。
📋 目录
  1. 先做离线复现,再谈 checkpoint 损坏
  2. 对比代码、数据与环境差异
  3. 检查 checkpoint 里实际保存了什么
  4. 多卡数据顺序与随机种子
  5. 常见问题
A A

模型训练时遇到「加载旧 checkpoint 后 loss 突然变大」,先别急着判断权重损坏。“半月前”这个时间差本身就是一条重要线索:这期间代码、数据、依赖甚至训练配置都可能改过。建议把排查拆成两个子问题:checkpoint 能否恢复到保存时的状态;当前环境是否还能按原来的方式组织样本和计算 loss。

如果同一环境下加载旧 checkpoint 后 loss 仍明显偏离保存时的水平,说明问题不在权重本身,而是数据、分布式采样或训练配置发生了变化;如果离线复现就对不上,则优先检查代码版本、预处理和依赖差异,而不是去调学习率或重新初始化。

先做离线复现,再谈 checkpoint 损坏

离线复现的意思是:不继续训练,只加载参数,用固定输入前向一次,比较 loss。如果这个 loss 和保存时记录的 loss 相当,说明模型参数没有被破坏;如果对不上,先不要急着检查多卡通信或学习率,而是要去看数据、代码和依赖。

ckpt = torch.load('model.pt', map_location='cpu')
model = build_model()
model.load_state_dict(ckpt['model'])  # key 名按实际保存结构调整
model.eval()

for i, batch in enumerate(fixed_loader):
    if i >= 10:
        break
    with torch.no_grad():
        out = model(**batch)
    print(i, out.loss.item())

fixed_loader 可以先从当前 Dataset 里按固定索引取出几个 batch 手动拼成一个 list,不走随机采样;如果保存过 global_step,直接用该 step 附近的输入做对比更准确。比较时注意模型模式:保存时是 train 就设 train,保存时是 eval 就设 eval,否则 dropout 和 BN 都会影响 loss。

对比代码、数据与环境差异

没有改代码的情况下,依赖或硬件变化也可能让同一个 checkpoint 在重载后 loss 有差异。以下检查项建议逐项确认,并记录到当前实验备注里:

检查项确认方式
PyTorch 版本torch.__version__
CUDA 版本torch.version.cuda
cuDNN 版本torch.backends.cudnn.version()
TF32 开关torch.backends.cuda.matmul.allow_tf32
cudnn benchmarktorch.backends.cudnn.benchmark
数据预处理 / collate / loss 计算git diff 或文件改动记录
总 batch size / 梯度累积步数训练日志或命令行参数

如果 checkpoint 里没有保存超参数,可以从训练日志或启动命令中找回。没有日志时优先检查 data/collate 和 forward 中 loss 的计算逻辑是否在半月内被改过。不同 PyTorch 版本对 TF32 和某些底层算子默认行为不一致,需要结合环境版本确认。

检查 checkpoint 里实际保存了什么

加载前先看 checkpoint 里面到底存了什么。很多加载报错或静默丢键都来自这里。常见的差异点是:只有 model state_dict,没有 optimizer 和 RNG;DDP 保存的 key 带 module. 前缀;AMP 训练时没有保存 scaler。下面的小脚本可以快速列出结构:

ckpt = torch.load('model.pt', map_location='cpu')
print('top keys:', list(ckpt.keys()))

for k in ckpt:
    if isinstance(ckpt[k], dict):
        print(k, len(ckpt[k]))
        print(list(ckpt[k].keys())[:5])

看到 torch_rng_state、cuda_rng_state、optimizer、scheduler、scaler 这些字段,说明保存内容较完整;只有模型字典时,后续恢复需要自己处理数据顺序和优化器状态。key 里带 module. 前缀时,单卡加载要去掉前缀,多卡加载则需要确认是否要保留,不能直接报错后忽略。

多卡场景中加载半月前checkpoint后loss突然变大的调试思路

多卡数据顺序与随机种子

多卡场景中,即使模型权重完全一致,各 rank 拿到的 batch 顺序也可能不同。数据加载线程的随机种子没有固定,是导致 loading 后 loss 偏高的常见原因。除 seed 外,还要看 batch size 与梯度累积步数是否与保存时一致;这两项改变会让同一步的 loss 在数值上不具可比性。

def seed_worker(worker_id):
    worker_seed = torch.initial_seed() % 2**32
    np.random.seed(worker_seed + worker_id)
    random.seed(worker_seed + worker_id)

def init_seed(seed, rank):
    random.seed(seed + rank)
    np.random.seed(seed + rank)
    torch.manual_seed(seed + rank)
    torch.cuda.manual_seed_all(seed + rank)

# DataLoader 中传入 worker_init_fn=seed_worker
# DistributedSampler 恢复时调用 set_epoch(旧 epoch)

如果旧 checkpoint 没有保存 rng_state,单步 loss 复现不出来很正常,可以对比几十步的滑动平均或变化趋势。恢复训练时如果发现 loss 起点整体抬高,先别调低学习率,否则会把加载错误和模型状态问题混在一起,越调越难定位。旧代码如果还能回滚,优先回滚 data、collate、loss 相关改动;不能回滚时再考虑从当前 checkpoint 继续训练,但要对可能出现偏差的模块单独验证。

常见问题

加载后 loss 是否一定要和原来一模一样?

不一定。训练模式下有 dropout、BN 统计更新,数据加载顺序变化都会让单步 loss 出现波动。建议对齐模型模式和 seed 后,对比几十步的均值或滑动平均;只有出现 NaN、Inf 或数量级跳变时,才优先怀疑加载错误。

单卡正常、多卡异常,优先排查什么?

先核对总 batch size、梯度累积步数、DataLoader worker seed 和 DistributedSampler 的 epoch 是否一致;再检查保存文件里是否有 DDP 的 module 前缀。用单卡加载和多卡加载做一次对照,能快速区分是权重加载问题还是多卡采样问题。

loss 变大后是否只能重新训练?

不建议直接重训。先从离线复现判断 checkpoint 是否损坏;再确认代码、数据和环境差异。如果旧代码不可回滚,可以尝试从 checkpoint 继续训练,但要接受 loss 起点变化,并单独验证可疑模块,而不是用调低学习率去掩盖加载错误。