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

如何保存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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 15:03:10