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

