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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 05:09:21