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

PyTorch下如何获取神经网络指定输出对参数phi的梯度并更新?

计算ag、bg对模型参数phi梯度的方案

首先给出修正缩进后的编码器PyTorch实现:

import torch
import torch.nn as nn

class Encoder(torch.nn.Module):
    def __init__(self, _l_dim, _hidden_dim, _fg_dim):
        super(Encoder, self).__init__()
        self.hidden_nn = nn.Linear(_l_dim, _hidden_dim)
        self.ag_nn = nn.Linear(_hidden_dim, _fg_dim)
        self.bg_nn = nn.Linear(_hidden_dim, _fg_dim)

    def forward(self, _lg):
        ag = self.ag_nn(self.hidden_nn(_lg))
        bg = self.bg_nn(self.hidden_nn(_lg))
        return ag, bg

基于PyTorch自动微分的实现方案

你可以直接用PyTorch内置的自动微分工具快速获取梯度,操作步骤如下:

  • 实例化模型后,所有可学习参数即为你提到的phi,可通过model.parameters()直接遍历
  • 前向传播得到ag、bg后,分别对两个输出做反向传播,注意首次反向时要保留计算图,避免计算bg梯度时计算图被释放
  • 反向传播后,每个参数的grad属性存储的就是对应梯度值

参考代码示例:

# 超参数示例
l_dim = 128
hidden_dim = 256
fg_dim = 64
batch_size = 8

# 实例化模型与构造输入
model = Encoder(l_dim, hidden_dim, fg_dim)
input_lg = torch.randn(batch_size, l_dim)
ag, bg = model(input_lg)

# 计算∂ag/∂phi
ag.sum().backward(retain_graph=True)
ag_grads = [param.grad.clone() for param in model.parameters()]
# 清空梯度缓存,避免和bg梯度累加
model.zero_grad()

# 计算∂bg/∂phi
bg.sum().backward()
bg_grads = [param.grad.clone() for param in model.parameters()]

如果你需要获取单个样本对应的梯度,不需要对batch维度做sum,单独对每个样本对应的输出元素反向即可。

手动梯度推导公式

你也可以根据网络结构手动推导梯度分量,先定义参数符号:

  • 隐含层hidden_nn权重为$W_h \in R^{hidden_dim \times l_dim}$,偏置为$b_h \in R^{hidden_dim}$
  • ag输出层ag_nn权重为$W_a \in R^{fg_dim \times hidden_dim}$,偏置为$b_a \in R^{fg_dim}$
  • bg输出层bg_nn权重为$W_b \in R^{fg_dim \times hidden_dim}$,偏置为$b_b \in R^{fg_dim}$
  • 输入为$x = _lg \in R^{batch \times l_dim}$,隐含层输出为$h = xW_h^T + b_h$

∂ag/∂phi各分量计算

  • 对$W_a$的梯度:单个样本下,ag第k维输出对$W_a$第k行的梯度为h向量,batch维度下按样本广播
  • 对$b_a$的梯度:对应输出维度梯度为1
  • 对$W_h$的梯度:$W_a^T$与输入x的外积,按batch维度广播
  • 对$b_h$的梯度:为$W_a^T$
  • 对$W_b、b_b$的梯度:全0,ag的计算不涉及这两个参数

∂bg/∂phi各分量计算

  • 对$W_b$的梯度:单个样本下,bg第k维输出对$W_b$第k行的梯度为h向量,batch维度下按样本广播
  • 对$b_b$的梯度:对应输出维度梯度为1
  • 对$W_h$的梯度:$W_b^T$与输入x的外积,按batch维度广播
  • 对$b_h$的梯度:为$W_b^T$
  • 对$W_a、b_a$的梯度:全0,bg的计算不涉及这两个参数

拿到对应梯度后,你可以手动用梯度下降规则更新参数,也可以传入PyTorch优化器完成更新。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 08:09:03