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

PyTorch中MLP近似函数时梯度损失的权重依赖保持问题

解决PyTorch中MLP梯度近似的依赖保留问题

你的核心问题是计算$\frac{\partial V(x)}{\partial x}$时,没有保留其与MLP权重的梯度依赖关系,导致损失无法反向传播更新模型参数。下面是具体的问题分析和修正方案:

原代码的问题点

  1. 输入张量未开启梯度追踪:x_tensor默认requires_grad=False,无法建立模型输出对输入的梯度与模型参数的关联。
  2. Jacobian计算未保留计算图:torch.autograd.functional.jacobian在输入无梯度追踪时,返回的张量会断开与模型参数的梯度连接,后续手动设置requires_grad_()也无法恢复原有计算图。

修正方案

要保留梯度对模型参数的依赖,需确保:

  • 输入张量开启梯度追踪;
  • 计算$\frac{\partial V(x)}{\partial x}$时保留计算图,让损失的反向传播能传递到模型权重。

修正后的完整代码

import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim


model = nn.Sequential(
    nn.Linear(2, 10),
    nn.ReLU(),
    nn.Linear(10, 10),
    nn.ReLU(),
    nn.Linear(10, 1)
)

model.float()
loss_fn = nn.MSELoss()  
optimizer = optim.Adam(model.parameters(), lr=0.01)  # 调低学习率避免震荡

# Generate Samples 
# V(x) = x^T P x
# grad V(x) = 2Px
P = np.matrix([[20.1892, -26.6218],[-26.6218, 38.0375]])
N_S = 10
N = N_S**2 # amount of samples
x_1 = np.linspace(-3,3,N_S)
x_2 = np.linspace(-3,3,N_S)
x = np.array([(a,b) for a in x_1 for b in x_2])
S = np.zeros((N,2))
for i in range(N):
    S[i,:]=2*P@x[i,:]
        

# training
epoch = 1

while epoch<1000:
    S_tensor = torch.from_numpy(S).float()
    x_tensor = torch.from_numpy(x).float()
    x_tensor.requires_grad_(True)  # 开启输入的梯度追踪

    # 计算模型输出V(x)
    V_x = model(x_tensor)
    # 计算V(x)对x的梯度,create_graph=True保留计算图
    grad_V_x = torch.autograd.grad(
        outputs=V_x,
        inputs=x_tensor,
        grad_outputs=torch.ones_like(V_x),
        create_graph=True,
        retain_graph=True
    )[0]

    loss = loss_fn(grad_V_x, S_tensor)

    optimizer.zero_grad()
    loss.backward()  # 此时损失梯度可正确传递到模型参数
    optimizer.step()

    if epoch % 50 == 0:
        print(f"epoch {epoch} loss {loss.item():.4f}")
    epoch += 1 

关键细节说明

  • x_tensor.requires_grad_(True):必须开启输入张量的梯度追踪,才能让后续的梯度计算关联到模型参数。
  • torch.autograd.grad的create_graph=True:这个参数会保留梯度计算的计算图,使得损失反向传播时,梯度能从grad_V_x传递到MLP的权重参数。
  • 调低学习率:原代码的lr=0.1过高,容易导致训练震荡,改为0.01后训练更稳定。

运行修正后的代码,你会看到损失持续下降,说明模型参数正在根据梯度样本正确更新。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 12:35:41