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

