PyTorch中神经网络雅可比/海森矩阵快速计算及jacrev性能疑问
PyTorch中jacrev计算雅可比慢于解析函数的原因及优化方案
核心原因:自动微分 vs 解析计算的本质差异
你的测试结果里jacrev+vmap比解析函数慢,核心是自动微分的固有开销和不必要的vmap调度成本:
- 自动微分的额外操作:
jacrev基于反向模式自动微分(Reverse-Mode AD),需要全程追踪计算图、存储中间张量,反向传播时还要遍历计算图节点逐个计算梯度。而你的解析函数df(x)=2*x是直接对张量做广播式element-wise运算,完全不需要计算图追踪,是PyTorch优化最充分的操作类型,效率自然碾压自动微分。 - 冗余的vmap调用:你的测试场景中,
jacrev(f)本身已经可以处理批量输入(输入a是10000x10000的张量,jacrev(f)(a)直接返回与df(a)一致的结果),额外套一层vmap会增加批量调度的冗余开销,进一步拖慢速度。 - 大张量场景的开销放大:10000x10000的超大张量会把自动微分的额外开销放大,而解析函数的广播运算刚好适配这种大张量的并行优化。
优化方案:修正用法+编译加速
1. 移除冗余的vmap调用
直接使用jacrev(f)(a)替代vmap(jacrev(f))(a),可以减少vmap的调度开销。修改后的测试代码:
from torch.func import jacrev import torch import time a = torch.rand(10000, 10000) def f(x): return (x ** 2).sum(-1) def df(x): return 2 * x t0 = time.time() b = df(a) t1 = time.time() c = jacrev(f)(a) # 移除冗余vmap t2 = time.time() assert torch.allclose(b, c) print(t1 - t0, t2 - t1)
这会让jacrev的运行速度有所提升。
2. 用torch.compile编译自动微分逻辑
PyTorch的torch.compile可以对自动微分的计算图做深度优化,消除冗余操作,大幅拉近与解析函数的速度差距。示例:
from torch.func import jacrev import torch import time a = torch.rand(10000, 10000) def f(x): return (x ** 2).sum(-1) def df(x): return 2 * x # 编译jacrev生成的雅可比函数 compiled_jac_f = torch.compile(jacrev(f)) t0 = time.time() b = df(a) t1 = time.time() c = compiled_jac_f(a) t2 = time.time() assert torch.allclose(b, c) print(t1 - t0, t2 - t1)
编译后的自动微分速度会接近解析函数的水平,同时保留自动微分的灵活性。
3. 选择合适的自动微分模式
- 如果你的函数输出维度远小于输入维度,使用
jacfwd(正向模式AD)会比jacrev更快,因为正向AD的时间复杂度是O(输出维度),反向AD是O(输入维度)。
关于海森矩阵的计算
手动推导神经网络的海森矩阵几乎不可行,torch.func.hessian或jacrev(jacrev(f))是更可行的方案,同样可以通过torch.compile来加速:
from torch.func import hessian import torch def f(x): return (x ** 2).sum(-1) # 编译海森计算函数 compiled_hess_f = torch.compile(hessian(f)) a = torch.rand(1000, 1000) # 海森矩阵维度大,建议用稍小的张量测试 hessian_matrix = compiled_hess_f(a)
编译后的海森计算速度会比原生自动微分快很多,足以应对大部分神经网络的二阶导数计算需求。
内容的提问来源于stack exchange,提问作者Frank Tian
相关产品推荐
相关产品推荐

