如何保存VS Code中Jupyter Notebook训练的模型及进度?
模型与训练进度保存方案
一、模型的正确保存方式
你当前用torch.save(netG, 'xxx.pt')直接保存模型对象是可行的,但更推荐保存模型的状态字典(state_dict)——这是PyTorch官方推荐的方式,它只保存模型的权重参数,文件更小,且能在不同环境、匹配的模型结构下灵活加载。
保存状态字典的代码:
# 保存生成器权重 torch.save(netG.state_dict(), 'C:/6th sem/TT(L)/Face detection/imgGen_state_dict.pt') # 保存判别器权重 torch.save(netD.state_dict(), 'C:/6th sem/TT(L)/Face detection/imgDis_state_dict.pt')
后续加载模型时,需先初始化对应模型结构,再加载权重:
# 初始化和训练时一致的模型结构 netG = Generator(...) # 替换为你实际的Generator类定义 netD = Discriminator(...) # 加载预训练权重 netG.load_state_dict(torch.load('C:/6th sem/TT(L)/Face detection/imgGen_state_dict.pt')) netD.load_state_dict(torch.load('C:/6th sem/TT(L)/Face detection/imgDis_state_dict.pt')) # 根据需求设置模式:推理用eval(),继续训练用train() netG.eval() # netG.train()
二、保存完整训练进度(断点续训)
若要彻底避免重新训练,仅保存模型权重不够,还需保存优化器状态、当前训练到的epoch/batch数、损失记录等信息,这样下次可从断点继续训练。
保存完整训练断点的代码:
# 整理需保存的训练状态 checkpoint = { 'epoch': current_epoch, # 当前训练到的epoch数 'batch_idx': current_batch, # 当前epoch内的batch序号 'netG_state_dict': netG.state_dict(), 'netD_state_dict': netD.state_dict(), 'optimizerG_state_dict': optimizerG.state_dict(), # 生成器的优化器 'optimizerD_state_dict': optimizerD.state_dict(), # 判别器的优化器 'losses': losses # 可选:保存已记录的损失值 } # 保存断点文件 torch.save(checkpoint, 'C:/6th sem/TT(L)/Face detection/training_checkpoint.pt')
加载断点继续训练的代码:
# 初始化模型和优化器 netG = Generator(...) netD = Discriminator(...) optimizerG = torch.optim.Adam(netG.parameters(), lr=0.0002) optimizerD = torch.optim.Adam(netD.parameters(), lr=0.0002) # 加载断点 checkpoint = torch.load('C:/6th sem/TT(L)/Face detection/training_checkpoint.pt') netG.load_state_dict(checkpoint['netG_state_dict']) netD.load_state_dict(checkpoint['netD_state_dict']) optimizerG.load_state_dict(checkpoint['optimizerG_state_dict']) optimizerD.load_state_dict(checkpoint['optimizerD_state_dict']) start_epoch = checkpoint['epoch'] start_batch = checkpoint['batch_idx'] losses = checkpoint.get('losses', []) # 从断点处继续训练循环 for epoch in range(start_epoch, total_epochs): # 若当前epoch未完成,从对应的batch开始 batch_start = start_batch if epoch == start_epoch else 0 for batch_idx in range(batch_start, total_batches): # 你的训练逻辑代码 ...
三、训练进度显示问题排查
显示(1150/1583)说明训练尚未完成(1583为总batch或epoch数),Jupyter输出可能存在延迟,可在训练循环中添加打印语句实时确认:
print(f"当前epoch: {epoch}, 当前batch: {batch_idx}/{total_batches}")
如果训练意外中断,用上述断点续训代码即可从上次停止的位置继续,无需重新训练。
四、你当前保存方式的有效性验证
你之前直接保存模型对象的方式可以完整保存权重,但跨版本或模型定义变更时可能加载失败。可立即测试加载是否有效:
# 加载之前保存的模型 loaded_netG = torch.load('C:/6th sem/TT(L)/Face detection/imgGen.pt') # 输入测试张量验证模型可用性 test_input = torch.randn(1, 100) # 替换为你的模型输入维度 output = loaded_netG(test_input) print(output.shape) # 若能正常输出维度,说明保存有效
内容的提问来源于stack exchange,提问作者Abhigyan Dutta
相关产品推荐
相关产品推荐

