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

如何在PyTorch中保留create_graph高效计算逐元素雅可比梯度

高效计算PyTorch中逐样本逐元素梯度(保留create_graph=True)

问题背景

当前使用torch.autograd.functional.jacobian计算标量模型输出相对于输入的逐样本、逐元素梯度,要求保留create_graph=True以支持后续对梯度的求导操作。现有实现会生成完整的(N, N, F)雅可比张量,再通过torch.einsum提取对角线部分得到(N, F)的目标梯度:

import torch
from torch.autograd.functional import jacobian

def method_jac_strict(inputs, forward_fn):
    # inputs: (N, F)
    # forward_fn: (N, F) -> (N, 1)
    # output: (N, F).

    # 计算完整雅可比矩阵: (N, 1, N, F)
    d = jacobian(forward_fn, inputs, create_graph=True, strict=True)
    d = d.squeeze()  # (N, N, F)

    # 提取对角线块(每个输出对自身输入样本的梯度): (N, F)
    d = torch.einsum('iif->if', d)
    return d

补充说明:模型可能包含BatchNorm等依赖批量的层,样本间并非完全独立,但仅需关注每个标量输出对自身输入样本的梯度,忽略跨样本依赖项。

优化方案

方法1:利用torch.autograd.functional.vjp批量计算(内存更高效)

通过向量-雅可比积(VJP),针对每个样本构造单位向量,直接批量计算目标梯度,避免生成冗余的N×N×F张量:

import torch
from torch.autograd.functional import vjp

def method_vjp_batch(inputs, forward_fn):
    N, F = inputs.shape
    # 构造单位矩阵作为VJP的输入向量,形状(N, N)
    v = torch.eye(N, device=inputs.device, dtype=inputs.dtype)
    # 计算VJP,直接得到每个样本输出对自身输入的梯度
    _, grads = vjp(forward_fn, inputs, v, create_graph=True)
    # 输出形状: (N, F)
    return grads

方法2:PyTorch 2.0+ 使用torch.func.vmap结合单样本梯度计算

利用vmap自动向量化单样本的梯度计算逻辑,既保留计算图,又避免跨样本的无效计算:

import torch
from torch.func import vmap, grad

def method_vmap_single(inputs, forward_fn):
    # 定义单样本的梯度计算函数:输入单个样本(F,),输出对应梯度(F,)
    def single_sample_grad(x):
        return grad(lambda x: forward_fn(x.unsqueeze(0)).squeeze())(x)
    
    # 使用vmap批量处理所有样本
    grads = vmap(single_sample_grad)(inputs)
    # 输出形状: (N, F)
    return grads

方法3:自定义autograd.Function(精细控制梯度传播)

如果需要更底层的计算控制,可以自定义autograd.Function,在反向传播时仅保留每个样本对自身输入的梯度贡献:

class PerSampleGrad(torch.autograd.Function):
    @staticmethod
    def forward(ctx, inputs, forward_fn):
        ctx.forward_fn = forward_fn
        ctx.save_for_backward(inputs)
        return forward_fn(inputs)
    
    @staticmethod
    def backward(ctx, grad_output):
        inputs, = ctx.saved_tensors
        N, _ = inputs.shape
        # 构造仅包含对角线的梯度权重,屏蔽跨样本梯度传播
        grad_output_diag = grad_output * torch.eye(N, device=inputs.device).unsqueeze(-1)
        # 计算梯度时仅保留自身样本的贡献
        grad_inputs = torch.autograd.grad(
            ctx.forward_fn(inputs), inputs, 
            grad_outputs=grad_output_diag, 
            create_graph=True,
            retain_graph=True
        )[0]
        return grad_inputs, None

def method_custom_autograd(inputs, forward_fn):
    return PerSampleGrad.apply(inputs, forward_fn)

方案对比

  • 原方法:实现简单,但大批次下(N, N, F)张量会占用大量内存,效率较低。
  • VJP方法:内存占用低,直接生成目标形状的梯度,适合大批次场景。
  • vmap方法:代码简洁,PyTorch 2.0+推荐,自动并行化单样本计算,效率高。
  • 自定义Autograd方法:灵活性最强,适合需要精细控制梯度传播逻辑的场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 04:52:10