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
相关产品推荐
相关产品推荐

