PyTorch中如何正确更新torch.nn.Parameter的部分参数
问题背景
针对torch.nn.Parameter定义的参数的部分更新需求,共测试3种实现方式,仅方案(2)可正常运行,其余两种均触发运行时错误:
- 方案(1):初始化阶段将Parameter赋值给普通张量的指定位置
import torch class NET(torch.nn.Module): def __init__(self): super(NET, self).__init__() self.params = torch.ones(4) self.P = torch.nn.Parameter(torch.ones(1)) self.params[1] = self.P def forward(self, x): y = x * self.params return y.sum() net = NET() x = torch.rand(4) optim = torch.optim.Adam(net.parameters(), lr=0.001) for _ in range(10): optim.zero_grad() loss = net(x) loss.backward() optim.step()
运行时报错:RuntimeError: Trying to backward through the graph a second time(尝试对计算图执行二次反向传播)
- 方案(2):前向传播阶段动态创建张量再赋值Parameter
import torch class NET(torch.nn.Module): def __init__(self): super(NET, self).__init__() self.P = torch.nn.Parameter(torch.ones(1)) def forward(self, x): params = torch.ones(4) params[1] = self.P y = x * params return y.sum() net = NET() x = torch.rand(4) optim = torch.optim.Adam(net.parameters(), lr=0.001) for _ in range(10): optim.zero_grad() loss = net(x) loss.backward() optim.step()
可正常运行,但每次前向传播都需要重新创建参数张量并完成赋值,存在额外性能开销
- 方案(3):初始化全量Parameter后修改指定位置的
requires_grad属性
import torch class NET(torch.nn.Module): def __init__(self): super(NET, self).__init__() self.params = torch.nn.Parameter(torch.ones(4)) def forward(self, x): y = x * self.params return y.sum() net = NET() net.params[1].requires_grad = False x = torch.rand(4) optim = torch.optim.Adam(net.parameters(), lr=0.001) for _ in range(10): optim.zero_grad() loss = net(x) loss.backward() optim.step()
运行时报错:RuntimeError: you can only change requires_grad flags of leaf variables.(仅可修改叶节点变量的requires_grad标记)
错误原因分析
- 方案(1)报错核心原因:
__init__阶段执行self.params[1] = self.P后,第一次前向传播生成的计算图节点会被持久保存在self.params属性中,后续迭代反向传播时会调用已经被释放的旧计算图,触发二次反向传播错误。本质是跨迭代保留了动态图上下文,违反了PyTorch动态图每轮前向重建的规则。 - 方案(3)报错核心原因:
nn.Parameter本身是叶节点张量,但对其做切片操作得到的net.params[1]是计算过程生成的非叶中间张量,PyTorch不允许修改非叶节点的requires_grad标记,因此直接报错。
低开销实现方案(规避两类报错)
参考方案(1)(3)的实现思路,不需要每次前向传播重建全量张量,有两种成熟实现可以满足部分参数更新的需求:
方案A:掩码屏蔽不需要更新的参数梯度(最灵活)
保留完整的nn.Parameter定义,反向传播结束后、优化器更新前,直接将不需要更新的位置的梯度置零,等价于对应位置参数不参与更新,完全规避非叶节点属性修改、计算图残留的问题,额外开销可忽略。
import torch class NET(torch.nn.Module): def __init__(self): super(NET, self).__init__() self.params = torch.nn.Parameter(torch.ones(4)) # 定义更新掩码:值为0的位置不更新,值为1的位置正常更新 self.update_mask = torch.tensor([1, 0, 1, 1], dtype=torch.bool) def forward(self, x): y = x * self.params return y.sum() net = NET() x = torch.rand(4) optim = torch.optim.Adam(net.parameters(), lr=0.001) for _ in range(10): optim.zero_grad() loss = net(x) loss.backward() # 屏蔽冻结位置的梯度 with torch.no_grad(): net.params.grad[~net.update_mask] = 0 optim.step()
该方案优势:
- 不需要在前向传播中动态创建张量,无重复初始化开销
- 代码改动量极小,支持训练过程中动态调整冻结/更新的位置
- 兼容所有优化器,不需要修改优化器参数组配置
方案B:拆分可训练/固定参数(性能最优)
将需要更新的参数单独定义为nn.Parameter,固定不变的参数注册为模型缓冲区(不会被优化器识别更新),前向传播时用torch.cat拼接得到完整参数。torch.cat是极轻量的张量操作,开销远低于每次重建全量张量,且不会留存跨迭代的计算图。
import torch class NET(torch.nn.Module): def __init__(self): super(NET, self).__init__() # 固定参数注册为缓冲区,不会被优化器更新 self.register_buffer('fixed_params', torch.tensor([1.0, 1.0, 1.0])) # 需要更新的参数单独定义为Parameter self.train_param = torch.nn.Parameter(torch.ones(1)) # 预定义拼接索引,避免前向传播重复计算 self.param_order = [0, 3, 1, 2] def forward(self, x): # 拼接得到完整参数,无计算图残留问题 params = torch.cat([self.fixed_params, self.train_param])[self.param_order] y = x * params return y.sum() net = NET() x = torch.rand(4) # 优化器自动识别可训练参数,训练循环不需要额外逻辑 optim = torch.optim.Adam(net.parameters(), lr=0.001) for _ in range(10): optim.zero_grad() loss = net(x) loss.backward() optim.step()
该方案优势:
- 训练循环不需要加额外梯度处理逻辑,训练阶段零额外开销
- 可训练参数会被
net.parameters()自动收集,不需要手动过滤 - 完全规避计算图残留问题,不会触发二次反向传播错误
方案选型建议
- 如果训练过程中需要动态切换参数的冻结/更新状态,优先选择方案A,仅需更新掩码即可完成调整
- 如果冻结/更新位置在训练前就已确定、不会动态变更,优先选择方案B,训练流程最简洁,性能最优
内容的提问来源于stack exchange,提问作者BONNED
相关产品推荐
相关产品推荐

