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

PyTorch:如何临时保存参数与梯度以便后续循环复用?

在PyTorch训练循环中内存式保存/加载参数、梯度与优化器状态

嘿,我完全懂你的需求——不想每次都写磁盘文件再读取,就在内存里完成每一步的状态快照和恢复对吧?其实PyTorch本身就提供了足够的工具来实现这个,不用碰任何.pt/.pth文件,直接用Python字典和张量的clone()方法就能轻松搞定。

核心思路

我们需要保存两类关键内容,缺一不可:

  • 模型的参数(权重)和对应的梯度:这部分可以通过遍历模型的named_parameters()来获取并克隆保存
  • 优化器的内部状态:比如SGD的动量缓冲、Adam的一阶/二阶矩估计等,这部分可以通过优化器的state_dict()来获取,注意要深拷贝其中的张量,避免后续训练更新污染保存的状态

具体实现代码示例

下面是一个完整的训练循环示例,每一步都会保存当前状态,并且演示如何在需要时快速恢复:

import torch
import torch.nn as nn
import torch.optim as optim

# 定义一个简单的测试模型
class SimpleModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(10, 5)
        self.fc2 = nn.Linear(5, 1)

model = SimpleModel()
optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
criterion = nn.MSELoss()

# 用字典存储每一步的状态,以循环步数为键
step_states = {}

# 模拟训练循环
for step in range(5):
    # 生成模拟训练数据
    x = torch.randn(32, 10)
    y = torch.randn(32, 1)
    
    # 标准训练流程
    optimizer.zero_grad()
    outputs = model(x)
    loss = criterion(outputs, y)
    loss.backward()
    
    # --- 关键:保存当前步的状态 ---
    # 1. 保存模型参数和梯度
    model_state = {}
    for name, param in model.named_parameters():
        # 用clone()创建独立副本,避免后续更新修改保存的快照
        model_state[name] = {
            'weight': param.data.clone(),
            'grad': param.grad.clone() if param.grad is not None else None
        }
    
    # 2. 保存优化器状态:逐个克隆状态里的张量,避免引用污染
    optimizer_state = {
        'param_groups': optimizer.state_dict()['param_groups'],
        'state': {}
    }
    for param_key, param_state in optimizer.state_dict()['state'].items():
        cloned_state = {}
        for k, v in param_state.items():
            if isinstance(v, torch.Tensor):
                cloned_state[k] = v.clone()
            else:
                cloned_state[k] = v
        optimizer_state['state'][param_key] = cloned_state
    
    # 将模型和优化器状态存入字典
    step_states[step] = {
        'model': model_state,
        'optimizer': optimizer_state
    }
    
    # 执行正常的参数更新
    optimizer.step()
    
    print(f"完成第 {step} 步训练,状态已保存至内存")

# --- 演示:恢复某一步的状态,比如恢复第2步的状态 ---
print("\n开始恢复第2步的训练状态...")
target_step = 2
saved_state = step_states[target_step]

# 1. 恢复模型参数和梯度
for name, param in model.named_parameters():
    saved_param = saved_state['model'][name]
    param.data.copy_(saved_param['weight'])
    if saved_param['grad'] is not None:
        param.grad = saved_param['grad'].clone()
    else:
        param.grad = None

# 2. 恢复优化器状态
optimizer.load_state_dict(saved_state['optimizer'])

# 验证:打印恢复后的模型fc1权重和梯度,和保存时完全一致
print("恢复后fc1权重:", model.fc1.weight.data)
print("恢复后fc1梯度:", model.fc1.grad)

关键细节提醒

  • 必须用clone():PyTorch的张量是引用类型,如果直接赋值,后续的参数更新会同步修改你保存的状态。clone()会创建一个独立的张量副本,确保保存的是当前步的静态快照。
  • 优化器状态的特殊处理:optimizer.state_dict()里的state字段包含了每个参数的优化器内部状态(比如动量的momentum_buffer),这些都是张量,必须逐个克隆,否则保存的是引用,后续训练会改变它。
  • 梯度可能为None:虽然backward()后梯度通常存在,但如果某个参数被冻结(requires_grad=False)或者没参与计算图,梯度会是None,所以保存和恢复时要处理这种情况。
  • 内存占用控制:如果训练步数很多,这个状态字典会占用较多内存,你可以根据需求只保留最近几步的状态,或者定期清理不需要的历史快照。

内容的提问来源于stack exchange,提问作者Daniel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 11:24:43