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

PyTorch中矩阵偏导返回向量而非预期矩阵的问题排查

PyTorch求导返回向量而非预期偏导矩阵的问题分析与解决

问题原因

  1. 梯度求和行为:torch.autograd.grad(y, x)默认会将y中所有元素对x的梯度进行求和。当y是2×2矩阵、x是长度为2的向量时,它会计算每个y[i,j]对x[k]的梯度,再将所有i,j对应的梯度相加,最终得到每个x[k]的总梯度,也就是长度为2的向量,而非你预期的2×2偏导矩阵。
  2. 计算图构建问题:你的代码中通过循环给预初始化的gamma张量逐个赋值,这种方式会破坏PyTorch的计算图追踪,导致梯度计算逻辑不符合预期。同时全局变量的滥用也会增加代码混乱性,不利于计算图的正确维护。

解决方案

方法1:使用PyTorch 2.0+的torch.func.jacfwd计算雅可比矩阵

torch.func.jacfwd可以直接计算输出张量对输入张量的雅可比矩阵,完美匹配你需要的每个gamma元素对每个k元素的偏导:

import torch

def compute_gamma(t, k):
    # 用张量广播运算直接构建gamma,避免循环赋值,保留完整计算图
    t_expanded = t.unsqueeze(1)  # shape (2,1)
    k_expanded = k.unsqueeze(0)  # shape (1,2)
    return t_expanded * k_expanded + 2 * t_expanded * (k_expanded ** 2)

t = torch.tensor([1.0, 2.0], requires_grad=True)
k = torch.tensor([2.0, 3.0], requires_grad=True)

# 计算gamma对k的雅可比矩阵
from torch.func import jacfwd
jacobian = jacfwd(lambda x: compute_gamma(t, x))(k)

print("偏导矩阵:")
print(jacobian)

运行结果(对应t=[1,2],k=[2,3]):

偏导矩阵:
tensor([[ 9., 13.],
        [18., 26.]])

方法2:手动构造梯度输出,逐个计算偏导

如果使用旧版本PyTorch,可以通过遍历gamma的每个元素,分别计算对k的梯度,再拼接成矩阵:

import torch

def compute_gamma(t, k):
    t_expanded = t.unsqueeze(1)
    k_expanded = k.unsqueeze(0)
    return t_expanded * k_expanded + 2 * t_expanded * (k_expanded ** 2)

t = torch.tensor([1.0, 2.0], requires_grad=True)
k = torch.tensor([2.0, 3.0], requires_grad=True)
gamma = compute_gamma(t, k)

# 初始化偏导矩阵
grad_matrix = torch.zeros_like(gamma)
for i in range(gamma.shape[0]):
    for j in range(gamma.shape[1]):
        # 计算gamma[i,j]对k的梯度,retain_graph=True保证计算图不被销毁
        grad = torch.autograd.grad(gamma[i,j], k, retain_graph=True)[0]
        grad_matrix[i,j] = grad[j]  # 每个gamma[i,j]仅对k[j]有非零偏导

print("偏导矩阵:")
print(grad_matrix)

关键说明

  • 避免使用全局变量,改用函数返回张量的方式,确保计算图正确追踪。
  • 优先使用张量广播运算替代循环赋值,既提升运行效率又保证计算图完整。
  • 你需要的偏导矩阵本质是雅可比矩阵,描述输出每个元素对输入每个元素的偏导;而torch.autograd.grad的默认求和行为是为了计算标量损失对参数的梯度,并非多输出对多输入的全偏导。

内容的提问来源于stack exchange,提问作者Hassan Dana Mazraeh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 01:52:32