PyTorch中如何保存/加载含多项损失的模型检查点
问题解答
是否必须存储分项损失才能恢复训练?
不需要。
训练中断后能否正常恢复,核心取决于是否保存了会影响后续参数更新的状态类数据:
- 模型权重参数(
model.state_dict()) - 优化器状态(含动量、各参数学习率等,
optimizer.state_dict()) - 学习率调度器状态(如果调度器按迭代/epoch自动调整学习率必须存储,你当前代码未保存该部分,续训时可能出现学习率异常)
- 当前训练进度(已完成的epoch数、batch步数)
你当前计算的分项损失、总损失都属于训练过程的日志指标,仅用于观测收敛情况,反向传播完成参数更新后,这些损失张量就不会再参与后续训练流程,哪怕完全不存任何损失值,只要上述核心状态保存完整,就能正常恢复训练。
如果需要在续训后连续记录各token特征的分项损失曲线、不中断训练日志统计,可以把分项损失作为日志字段存入检查点,这部分数据体积很小,不会带来额外存储负担。
具体实现方法
第一步:修改训练循环逻辑,计算epoch维度的平均分项损失
注意不要直接保存带计算图关联的PyTorch张量,必须调用.item()转换为普通Python数值,避免检查点携带无用计算图导致体积膨胀、加载报错。
修改后的训练循环参考代码:
def train(epoch_num, model, dataloader, loss_func, opt, lr_scheduler, num_iters=-1): # 初始化累计值 total_loss_sum = 0.0 per_class_loss_sum = [0.0 for _ in range(LEN_VOCAB)] sample_num = 0 for batch_num, batch in enumerate(dataloader): opt.zero_grad() x = batch[0].to(get_device()) tgt = batch[1].to(get_device()) batch_size = x.shape[0] y = model(x.permute(1, 0, 2)) losses = [] for j in range(LEN_VOCAB): aux_loss = loss_func.forward(y[j].permute(1, 2, 0), tgt[..., j]) losses.append(aux_loss) losses_sum = sum(losses) losses_sum.backward() opt.step() if lr_scheduler is not None: lr_scheduler.step() lr = opt.param_groups[0]['lr'] # 累加损失,转纯数值脱离计算图 total_loss_sum += losses_sum.item() * batch_size for j in range(LEN_VOCAB): per_class_loss_sum[j] += losses[j].item() * batch_size sample_num += batch_size if batch_num == num_iters: break # 计算整个epoch的平均损失 avg_total_loss = total_loss_sum / sample_num avg_per_class_loss = [loss_sum / sample_num for loss_sum in per_class_loss_sum] # 同时返回总损失和分项损失 return avg_total_loss, avg_per_class_loss
第二步:修改epoch循环,将分项损失存入检查点
建议同时补充保存学习率调度器状态,避免续训时学习率异常:
for epoch in range(0, epochs): print('Epoch: ', epoch) # 接收训练函数返回的总损失和分项损失 total_loss, per_class_loss = trfrmr.train(epoch+1, model, train_loader, train_loss_func, opt, lr_scheduler, num_iters=-1) loss_train.append(total_loss) # 如有单独存分项损失的日志列表,可同步append # per_class_loss_train.append(per_class_loss) torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': opt.state_dict(), 'lr_scheduler_state_dict': lr_scheduler.state_dict() if lr_scheduler is not None else None, # 新增保存调度器状态 'total_loss': total_loss, 'per_class_loss': per_class_loss, # 新增保存分项损失 }, "model_pop909_checkpoint.pth")
第三步:续训加载检查点的对应逻辑
加载时对应读取字段即可,哪怕不读取per_class_loss字段也完全不影响训练正常运行:
checkpoint = torch.load("model_pop909_checkpoint.pth", map_location=get_device()) start_epoch = checkpoint['epoch'] + 1 model.load_state_dict(checkpoint['model_state_dict']) opt.load_state_dict(checkpoint['optimizer_state_dict']) if checkpoint['lr_scheduler_state_dict'] is not None and lr_scheduler is not None: lr_scheduler.load_state_dict(checkpoint['lr_scheduler_state_dict']) # 如需续接日志,读取历史损失即可 # loss_train = [checkpoint['total_loss']] # per_class_loss_train = [checkpoint['per_class_loss']]
注意事项
- 所有存入检查点的损失值必须是调用
.item()得到的普通数值,不能直接存反向传播用的损失张量,否则会导致检查点体积异常增大。 - 如果你使用的学习率调度器是按epoch更新(比如
StepLR),除了保存调度器状态,还要注意调度器的step()调用位置要和训练逻辑匹配,避免学习率计算错误。
内容的提问来源于stack exchange,提问作者Enrique Vilchez Campillejo
相关产品推荐
相关产品推荐

