PyTorch深度学习模型续训:模型保存与加载问题咨询
嘿,我完全懂你这种续训卡壳的烦恼——之前我也踩过只存模型state_dict的坑!问题出在你漏掉了优化器的状态信息,这对续训来说至关重要。下面给你一步步拆解解决方案:
一、为什么只存模型参数不够?
PyTorch的model.state_dict()只保存模型的权重参数,但续训时不仅需要模型参数,还得保留优化器的状态(比如Adam的动量项、累计梯度等)。如果只加载模型参数,相当于用全新的优化器重新训练,自然达不到“接着上次进度训练”的效果。
二、需要保存的核心内容
你需要把以下内容打包成一个字典保存:
- 模型的
state_dict - 优化器的
state_dict - 训练元数据(当前训练到的epoch数、最新损失值,方便后续追踪)
三、合适的保存时机
建议在每日训练结束后(或者每个epoch结束后)保存一次checkpoint。如果是每日固定续训,直接覆盖同一个checkpoint文件即可;如果需要保留历史版本,可以按日期命名(比如checkpoint_20240520.pt)。
四、针对你代码的具体修改
1. 修改训练函数,添加保存逻辑
在你的trainv函数里,每个epoch循环结束后加入保存代码:
def trainv(model, device, epochs, train_iterator, optimizer, validate_iterator, start_epoch=0): train_losses = [] validate_losses = [] for epoch in range(start_epoch, start_epoch + epochs): model.train() epoch_train_loss = 0.0 # 批量训练 for local_batch, local_labels in train_iterator: train_loss = train_batch(model, optimizer, device) epoch_train_loss += train_loss.item() * local_batch.size(0) epoch_train_loss /= len(train_iterator.dataset) train_losses.append(epoch_train_loss) # 验证环节 validate_loss = runs_for_validate(validate_iterator, n_samples) validate_losses.append(validate_loss) # 保存每日训练后的checkpoint checkpoint = { 'epoch': epoch + 1, # 记录下一次要开始的epoch 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'latest_train_loss': epoch_train_loss, 'latest_val_loss': validate_loss } torch.save(checkpoint, 'daily_checkpoint.pt') # 覆盖旧的checkpoint return train_losses, validate_losses
2. 加载checkpoint续训的代码
第二天训练时,先加载之前的checkpoint,再继续训练3-4个epoch:
# 初始化模型和优化器(必须和上次训练的参数完全一致!) model = NET(inputs).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate) # 加载保存的checkpoint # 如果保存时用GPU,现在用CPU训练的话,要加map_location=device checkpoint = torch.load('daily_checkpoint.pt', map_location=device) model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) start_epoch = checkpoint['epoch'] # 获取上次训练到的最后一个epoch+1 # 继续训练3-4个epoch train_losses, validate_losses = trainv( model, device, epochs=4, train_iterator=train_iterator, optimizer=optimizer, validate_iterator=validate_iterator, start_epoch=start_epoch )
五、常见坑点排查
- 加载时的模型结构、优化器参数(比如学习率)必须和保存时完全一致,否则会加载失败
- 如果用到了学习率调度器(比如
torch.optim.lr_scheduler),也要把调度器的state_dict加入checkpoint一起保存和加载 - 注意设备匹配:如果保存模型时用的是GPU,加载时如果切换到CPU,一定要加上
map_location=device参数
内容的提问来源于stack exchange,提问作者Sadcow
相关产品推荐
相关产品推荐

