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

PyTorch中如何求解参数空间线性组合权重λ的梯度?

问题根源分析

你代码里para_b.data = para_a.data * self.lambda_这一步直接操作了张量的.data属性,这会绕过PyTorch的自动微分系统——.data获取的是张量的底层原始数据,修改它不会被记录到计算图中,导致后续模型B的输出和λ之间没有建立梯度传播的关联,最终lambda_.grad为None。

解决方案

要让λ的梯度能被正确计算,必须让λ参与的运算过程被PyTorch的计算图完整记录,下面提供两种可行方案:

方案一:动态使用缩放后的参数计算(无需维护模型B参数)

直接在forward过程中复用模型A的参数,动态乘以λ后进行计算,不需要修改模型B的参数,这样λ的参与会被完整记录到计算图中。

修改后的代码:

import torch
import torch.nn as nn

class MyBaseModel(nn.Module):
    def __init__(self):
        super(MyBaseModel, self).__init__()
        self.linear1 = nn.Linear(3, 8)
        self.act1 = nn.ReLU()
        self.linear2 = nn.Linear(8, 4)
        self.act2 = nn.Sigmoid()
        self.linear3 = nn.Linear(4, 5)
    
    def forward(self, x, params=None):
        # 优先使用传入的缩放参数,否则用自身默认参数
        if params is None:
            linear1_w, linear1_b, linear2_w, linear2_b, linear3_w, linear3_b = (
                self.linear1.weight, self.linear1.bias,
                self.linear2.weight, self.linear2.bias,
                self.linear3.weight, self.linear3.bias
            )
        else:
            linear1_w, linear1_b, linear2_w, linear2_b, linear3_w, linear3_b = params
        
        x = torch.nn.functional.linear(x, linear1_w, linear1_b)
        x = self.act1(x)
        x = torch.nn.functional.linear(x, linear2_w, linear2_b)
        x = self.act2(x)
        x = torch.nn.functional.linear(x, linear3_w, linear3_b)
        return x

class WeightedSumModel(nn.Module):
    def __init__(self):
        super(WeightedSumModel, self).__init__()
        self.lambda_ = nn.Parameter(torch.tensor(2.0))
        self.a = MyBaseModel()
    
    def forward(self, x):
        # 对模型A的所有参数进行缩放
        scaled_params = [p * self.lambda_ for p in self.a.parameters()]
        # 用缩放后的参数计算输出
        return self.a(x, params=scaled_params).sum()

input_tensor = torch.ones((2, 3))
weighted_sum_model = WeightedSumModel()
output_tensor = weighted_sum_model(input_tensor)
output_tensor.backward()

print(weighted_sum_model.lambda_.grad)

方案二:使用参数重参数化(保留模型B结构)

如果需要保留模型B的结构,可以用PyTorch的parametrize工具给模型B的参数添加动态依赖,让它始终等于模型A参数乘以λ,这样自动微分系统会追踪λ的梯度。

修改后的代码:

import torch
import torch.nn as nn
from torch.nn.utils.parametrize import register_parametrization

class MyBaseModel(nn.Module):
    def __init__(self):
        super(MyBaseModel, self).__init__()
        self.linear1 = nn.Linear(3, 8)
        self.act1 = nn.ReLU()
        self.linear2 = nn.Linear(8, 4)
        self.act2 = nn.Sigmoid()
        self.linear3 = nn.Linear(4, 5)
    
    def forward(self, x):
        return self.linear3(self.act2(self.linear2(self.act1(self.linear1(x)))))

# 定义重参数化逻辑:将基础参数乘以λ
class ScaleParam(nn.Module):
    def __init__(self, lambda_param, base_param):
        super().__init__()
        self.lambda_ = lambda_param
        # 将基础参数注册为缓冲区,不参与优化
        self.register_buffer('base_param', base_param.detach().clone())
    
    def forward(self, _):
        # _是占位参数,实际使用缓冲区的基础参数乘以λ
        return self.base_param * self.lambda_

class WeightedSumModel(nn.Module):
    def __init__(self):
        super(WeightedSumModel, self).__init__()
        self.lambda_ = nn.Parameter(torch.tensor(2.0))
        self.a = MyBaseModel()
        self.b = MyBaseModel()
        
        # 给模型B的每个参数注册重参数化
        for (name_a, param_a), (name_b, param_b) in zip(self.a.named_parameters(), self.b.named_parameters()):
            # 替换模型B的原参数为占位张量
            setattr(self.b, name_b, nn.Parameter(torch.zeros_like(param_b)))
            # 注册重参数化规则
            register_parametrization(self.b, name_b, ScaleParam(self.lambda_, param_a))
    
    def forward(self, x):
        return self.b(x).sum()

input_tensor = torch.ones((2, 3))
weighted_sum_model = WeightedSumModel()
output_tensor = weighted_sum_model(input_tensor)
output_tensor.backward()

print(weighted_sum_model.lambda_.grad)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 09:17:08