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

如何用torch.autograd计算张量值函数的雅可比矩阵?含散度优化方案

问题解决方法

一、解决dy_dx为None的问题

你遇到的核心问题是计算y后覆盖x变量,导致传入jacobian的x与y的计算图依赖断裂:y是基于原始形状的x计算的,但你重新赋值x为扁平化后的tensor,此时新x并非y计算图中的叶子节点,autograd无法追踪到依赖关系,因此返回None。

修复方案(二选一)

方案1:保留原始x的叶子节点引用

修改调用逻辑,不覆盖原始x,直接基于原始x计算y,再分别扁平化y和原始x,同时调整jacobian函数对原始x求导后扁平化梯度:

def jacobian(y_flat, x_original):
    k, D = y_flat.shape
    jacobian = []
    for i in range(D):
        v = torch.zeros_like(y_flat)
        v[:, i] = 1.
        # 对原始形状的x求导,得到对应梯度后扁平化
        dy_dx = torch.autograd.grad(y_flat, x_original, grad_outputs=v, retain_graph=True, create_graph=True, allow_unused=True)[0]
        dy_dx_flat = dy_dx.flatten(1)
        jacobian.append(dy_dx_flat)
    jacobian = torch.stack(jacobian, dim=1).requires_grad_()
    return jacobian

# 调用示例
x_original = x  # 保留叶子节点引用
y = g(t, x_original)
y_flat = y.flatten(1)
jac = jacobian(y_flat, x_original)

方案2:基于扁平化后的x构建计算图

先将x扁平化,再reshape为g需要的形状传入,确保y的计算图依赖于扁平化后的叶子节点x:

# 扁平化x并保留为叶子节点
x_flat = x.flatten(1)
# 将扁平化后的x还原为g需要的形状(视图操作,不破坏梯度依赖)
x_reshaped = x_flat.view(x.shape)
y = g(t, x_reshaped)
y_flat = y.flatten(1)
# 直接调用原jacobian函数
jac = jacobian(y_flat, x_flat)

二、计算散度的高效方案

散度是雅可比矩阵的迹,无需计算完整雅可比矩阵,可通过两种方式优化:

方法1:循环取对角线元素求和(内存友好)

仅计算雅可比矩阵的对角线元素并求和,避免存储完整的D×D矩阵:

def divergence(y, x):
    k = y.shape[0]
    D = y.flatten(1).shape[1]
    div = torch.zeros(k, device=y.device)
    # 扁平化y和x,方便按索引操作
    y_flat = y.flatten(1)
    x_flat = x.flatten(1)
    for i in range(D):
        v = torch.zeros_like(y_flat)
        v[:, i] = 1.
        # 求y_i对x的梯度,取对应x_i的分量
        grad = torch.autograd.grad(y_flat, x_flat, grad_outputs=v, retain_graph=True, create_graph=True)[0]
        div += grad[:, i]
    return div

# 调用示例
y = g(t, x)
div = divergence(y, x)

方法2:利用自动求导特性的一次性计算(速度更快,内存占用较高)

当D较小时,可通过构造单位矩阵的grad_outputs一次性获取雅可比矩阵,再取对角线求和:

def divergence_fast(y, x):
    y_flat = y.flatten(1)
    x_flat = x.flatten(1)
    k, D = y_flat.shape
    # 构造(k, D, D)的单位矩阵作为grad_outputs
    grad_outputs = torch.eye(D, device=y.device).unsqueeze(0).repeat(k, 1, 1)
    # 计算完整雅可比矩阵
    jac = torch.autograd.grad(y_flat, x_flat, grad_outputs=grad_outputs, create_graph=True)[0]
    # 取对角线并求和得到散度
    div = jac.diagonal(dim1=1, dim2=2).sum(dim=1)
    return div

# 调用示例
y = g(t, x)
div = divergence_fast(y, x)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 10:32:33