JAX 前向与反向自动微分完全指南:JVP、VJP 与 Hessian-vector products 的原理与实战
2026/9/10 15:40:28 网站建设 项目流程

JAX 前向与反向自动微分完全指南:JVP、VJP 与 Hessian-vector products 的原理与实战

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

JAX 同时内置了前向模式(forward-mode)与反向模式(reverse-mode)两种自动微分(automatic differentiation, AD)实现,本文以官方进阶指南 docs/jacobian-vector-products.md 为主体,系统讲解 Jacobian-vector products(JVP)与 Vector-Jacobian products(VJP)的数学定义、类型签名、计算复杂度,以及如何用它们组合出 Hessian-vector products、矩阵-Jacobian 乘积,并最终还原jax.jacfwd/jax.jacrev的底层实现。读完本文,你将理解jax.grad为什么能高效训练百万、十亿级参数的神经网络,掌握jax.jvpjax.vjpjax.jacfwdjax.jacrevjax.hessian的选型逻辑,并能用jax.vmap组合它们写出高性能的自定义微分算子。

一、两种自动微分模式:为什么需要一点数学背景

熟悉的jax.grad构建在反向模式之上,但要讲清楚两种模式的差异、以及各自在什么场景下更有用,需要先铺垫一些数学背景。核心问题可以概括为:给定函数 $f : \mathbb{R}^n \to \mathbb{R}^m$,我们想要"沿着某个方向求导数"。

  • 前向模式直接计算 Jacobian-vector product(JVP),即 $\partial f(x) v$,代价约为一次函数求值的 3 倍,且内存与计算深度无关
  • 反向模式计算 vector-Jacobian product(VJP),即 $v^\mathsf{T} \partial f(x)$,一次调用就能得到标量损失函数的梯度,但内存随计算深度线性增长

JAX 对两种模式都有高效且通用的实现。相关 API 的完整定义集中在 jax/_src/api.py,底层 AD 解释器实现在 jax/_src/interpreters/ad.py,本文后续会逐一给出源码级证据。

二、前向模式:Jacobian-vector products(JVP)

2.1 数学定义:从 Jacobian 矩阵到 pushforward 线性映射

给定函数 $f : \mathbb{R}^n \to \mathbb{R}^m$,$f$ 在输入点 $x \in \mathbb{R}^n$ 处的 Jacobian 记为 $\partial f(x)$,通常被看作一个 $m \times n$ 矩阵:

$\qquad \partial f(x) \in \mathbb{R}^{m \times n}$。

但也可以把 $\partial f(x)$ 看作一个线性映射:它把 $f$ 的定义域在 $x$ 处的切空间(即另一份 $\mathbb{R}^n$)映到 $f$ 的值域在 $f(x)$ 处的切空间(即一份 $\mathbb{R}^m$):

$\qquad \partial f(x) : \mathbb{R}^n \to \mathbb{R}^m$。

这个映射被称为 $f$ 在 $x$ 处的 pushforward(前推)映射,Jacobian 矩阵只是该线性映射在标准基下的矩阵表示。

如果不固定具体的输入点 $x$,可以把 $\partial f$ 看成一个先接收输入点、再返回该点处 Jacobian 线性映射的函数:

$\qquad \partial f : \mathbb{R}^n \to \mathbb{R}^n \to \mathbb{R}^m$。

对输入点 $x \in \mathbb{R}^n$ 和一个切向量 $v \in \mathbb{R}^n$,我们得到输出切向量 $\in \mathbb{R}^m$。这个从 $(x, v)$ 对到输出切向量的映射,就是Jacobian-vector product

$\qquad (x, v) \mapsto \partial f(x) v$。

2.2 JAX 代码中的 JVP:jax.jvp

回到 Python 代码,JAX 的jax.jvp正是对这个变换的建模:给定一个计算 $f$ 的 Python 函数,jax.jvp返回一个计算 $(x, v) \mapsto (f(x), \partial f(x) v)$ 的 Python 函数。下面用官方指南中的玩具模型(sigmoid 二分类预测)演示:

import jax import jax.numpy as jnp key = jax.random.key(0) # 初始化随机模型系数 key, W_key, b_key = jax.random.split(key, 3) W = jax.random.normal(W_key, (3,)) b = jax.random.normal(b_key, ()) # 定义 sigmoid 函数 def sigmoid(x): return 0.5 * (jnp.tanh(x / 2) + 1) # 输出标签为真的概率 def predict(W, b, inputs): return sigmoid(jnp.dot(inputs, W) + b) # 构造玩具数据集 inputs = jnp.array([[0.52, 1.12, 0.77], [0.88, -1.08, 0.15], [0.52, 0.06, -1.30], [0.74, -2.49, 1.39]]) # 隔离出"从权重矩阵到预测值"的函数 f = lambda W: predict(W, b, inputs) key, subkey = jax.random.split(key) v = jax.random.normal(subkey, W.shape) # 沿 f 在 W 处前推向量 v y, u = jax.jvp(f, (W,), (v,))

这里primalstangents都要求是 tuple 或 list(jax/_src/api.py 中会强制校验类型,并检查二者的树结构与形状、dtype 是否匹配),返回值是一个(primals_out, tangents_out)对:y = f(W)u = ∂f(W)·v

从源码看,jax.jvp实际调用 jax/_src/interpreters/ad.py 中的ad.jvp:它会创建JVPTrace,把每个 primal 值配上对应的 tangent 值封装成JVPTracer重新执行fun。每遇到一个原始数值运算,就同时执行该运算的 "JVP 规则"——既在 primal 上求值,也在该 primal 点应用其 JVP——这正是下文复杂度结论的实现基础。

2.3 类型签名视角

借用 Haskell 风格的签名可以更精确地描述:

jvp :: (a -> b) -> a -> T a -> (b, T b)

其中T a表示a的切空间类型。也就是说,jvp接收:一个类型为a -> b的函数、一个类型为a的值、一个类型为T a的切向量;返回一个由类型为b的值和类型为T b的输出切向量组成的对。

jvp变换后的函数,其求值方式与原函数几乎一样,只是每个类型为a的 primal 值旁边都带上了类型为T a的 tangent 值。原函数每应用一个基本数值运算,jvp变换后的函数就执行该基本运算的 JVP 规则:既在 primal 上求值该运算,又在该 primal 值处应用该运算的 JVP。

2.4 计算复杂度:约 3 倍 FLOPs 与"与深度无关"的内存

这种边算边推的求值策略直接决定了复杂度特征:

  • 内存与计算深度无关:由于 JVP 是随求值过程即时推进的,不需要为后续存储任何中间结果,因此内存成本不随计算的深度增长;
  • FLOPs 约为原函数的 3 倍:一部分工作量用于求值原函数(例如sin(x)),一部分用于线性化(例如cos(x)),一部分用于把线性化函数作用到向量上(例如cos_x * v)。

换句话说,固定 primal 点 $x$ 后,以约等于一次f求值的边际代价,就可以对任意方向 $v$ 计算 $v \mapsto \partial f(x) \cdot v$。

2.5 为什么机器学习中很少单独使用前向模式

内存优势听起来很有吸引力,那为什么前向模式在机器学习里不常见?

关键在于如何用 JVP 拼出完整 Jacobian 矩阵:如果把 JVP 作用在 one-hot 切向量上,它就揭示 Jacobian 矩阵中与该非零分量对应的一列。因此可以逐列构建完整的 Jacobian,而每一列的成本都约等于一次函数求值。这对"高瘦"(tall)的 Jacobian 高效,但对"宽扁"(wide)的 Jacobian 低效。

而基于梯度的机器学习优化,目标函数是从参数空间 $\mathbb{R}^n$ 到标量损失 $\mathbb{R}$ 的映射,其 Jacobian 是一个极宽的矩阵:$\partial f(x) \in \mathbb{R}^{1 \times n}$,通常与梯度向量 $\nabla f(x) \in \mathbb{R}^n$ 等同。逐列构建这个矩阵、每列都花一次函数求值的 FLOPs,显然低效——尤其当 $f$ 是训练损失函数、$n$ 高达百万甚至十亿时,这种方法根本无法扩展。要做得更好,就需要反向模式。

三、反向模式:Vector-Jacobian products(VJP)

前向模式给出计算 Jacobian-vector product 的函数,可逐列构建 Jacobian;反向模式则给出计算 vector-Jacobian product(等价地,Jacobian 转置-向量乘积)的函数,可逐行构建 Jacobian。

3.1 数学定义:pullback 与转置

仍考虑函数 $f : \mathbb{R}^n \to \mathbb{R}^m$。沿用 JVP 的记号,VJP 的记号非常简洁:

$\qquad (x, v) \mapsto v \partial f(x)$,

其中 $v$ 是 $f$ 在 $x$ 处余切空间(cotangent space,与另一份 $\mathbb{R}^m$ 同构)中的元素。严格地说,应把 $v$ 看作线性映射 $v : \mathbb{R}^m \to \mathbb{R}$,把 $v \partial f(x)$ 理解为函数复合 $v \circ \partial f(x)$——类型之所以成立,是因为 $\partial f(x) : \mathbb{R}^n \to \mathbb{R}^m$。但在常见情况下可以把 $v$ 等同于 $\mathbb{R}^m$ 中的向量,两者几乎可以互换使用,就像我们有时会在"列向量"和"行向量"之间切换而不加说明一样。

有了这个等同,也可以把 VJP 的线性部分看作 JVP 线性部分的转置(或伴随共轭):

$\qquad (x, v) \mapsto \partial f(x)^\mathsf{T} v$。

对给定点 $x$,签名可以写作

$\qquad \partial f(x)^\mathsf{T} : \mathbb{R}^m \to \mathbb{R}^n$。

余切空间上的这个对应映射常被称为 $f$ 在 $x$ 处的 pullback(拉回)。对我们的目的而言,关键在于它从"看起来像 $f$ 输出"的量出发,得到"看起来像 $f$ 输入"的量——正如我们对转置线性映射的预期。

3.2 JAX 代码中的 VJP:jax.vjp

从数学回到 Python:JAX 的vjp接收一个计算 $f$ 的 Python 函数,返回一个计算 $(x, v) \mapsto (f(x), v^\mathsf{T} \partial f(x))$ 的 Python 函数:

from jax import vjp # 隔离出"从权重矩阵到预测值"的函数 f = lambda W: predict(W, b, inputs) y, vjp_fun = vjp(f, W) key, subkey = jax.random.split(key) u = jax.random.normal(subkey, y.shape) # 沿 f 在 W 处拉回余向量 u v = vjp_fun(u)

vjp返回的vjp_fun是一个线性函数:它接收与输出同形状的余切向量(cotangent),返回与每个输入同形状的余切向量。值得注意的是,vjp在这里是"分两步走"的——先调用vjp(f, W)完成正向求值与线性化、拿到vjp_fun,之后再多次调用vjp_fun(u)传入不同的余切向量做拉回。

从源码看(jax/_src/api.py),vjp通过ad.linearize对函数做线性化,生成切线方向的 jaxpr 与残差(residuals),再封装出可调用的VJP对象;其 docstring 明确指出jax.gradjax.vjp的特例("gradis implemented as a special case ofvjp")。

3.3 类型签名视角

同样可以用 Haskell 风格签名表示:

vjp :: (a -> b) -> a -> (b, CT b -> CT a)

其中CT a表示a的余切空间类型。vjp接收一个类型为a -> b的函数和一个类型为a的点,返回一个由类型为b的值和类型为CT b -> CT a的线性映射组成的对。

3.4 复杂度与jax.grad的效率来源

VJP 让我们可以逐行构建 Jacobian 矩阵,而计算 $(x, v) \mapsto (f(x), v^\mathsf{T} \partial f(x))$ 的 FLOPs 代价同样只有求值 $f$ 的约 3 倍。特别是,如果要求 $f : \mathbb{R}^n \to \mathbb{R}$ 的梯度,只需一次调用即可完成。这就是jax.grad对基于梯度的优化高效的原因——即使目标函数是参数以百万、十亿计的神经网络训练损失。

不过反向模式也有代价:虽然 FLOPs 友好,但内存随计算深度增长(需要保存正向传播中的中间值供反向使用),且其实现传统上比前向模式复杂得多——不过 JAX 对此有一些技巧,例如通过部分求值(partial evaluation)把线性化后的计算压缩成 tangent jaxpr 再执行反向传播,相关机制可参见 jax/_src/interpreters/ad.py 中的linearize/linearize_jaxpr/backward_pass系列函数。

3.5 用 VJP 实现向量值梯度

如果你需要的是向量值梯度(类似tf.gradients),可以用 VJP 这样实现:

def vgrad(f, x): y, vjp_fn = jax.vjp(f, x) return vjp_fn(jnp.ones(y.shape))[0] print(vgrad(lambda x: 3*x**2, jnp.ones((2, 2))))

其思路是:把全 1 向量作为输出余切传入vjp_fn,一次性拉回得到对输入每个分量的"梯度"。

四、Hessian-vector products:前向与反向的协同

4.1 纯反向模式的 HVP 基线

在前一节的基础上,先用纯反向模式实现一个 Hessian-vector product(假设二阶导数连续):

def hvp(f, x, v): return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)

这个实现是高效的,但还可以做得更好——把前向模式和反向模式组合起来,能进一步节省内存。

4.2 数学推导:对梯度做 JVP

给定要微分的函数 $f : \mathbb{R}^n \to \mathbb{R}$、线性化点 $x \in \mathbb{R}^n$ 和向量 $v \in \mathbb{R}^n$,我们想要的 Hessian-vector product 是:

$(x, v) \mapsto \partial^2 f(x) v$。

考虑辅助函数 $g : \mathbb{R}^n \to \mathbb{R}^n$,它是 $f$ 的导数(梯度),即 $g(x) = \partial f(x)$。我们只需要它的 JVP,因为:

$(x, v) \mapsto \partial g(x) v = \partial^2 f(x) v$。

这个推导几乎可以一字不差地翻译成代码:

# forward-over-reverse def hvp(f, primals, tangents): return jax.jvp(jax.grad(f), primals, tangents)[1]

这里jax.grad(f)是反向模式(外层是前向的jax.jvp),因此这种写法被称为forward-over-reverse(前向套反向)。更妙的是,由于不需要直接调用jnp.dot,这个hvp函数:

  • 适用于任意形状的数组;
  • 适用于任意容器类型(如嵌套的 list / dict / tuple 存储的向量);
  • 甚至不依赖jax.numpy模块。

4.3 用jax.jacfwd(jax.jacrev)验证正确性

下面用官方指南的示例验证:用jax.hessian的朴素实现(物化完整 Hessian 张量)与hvp对比。

def f(X): return jnp.sum(jnp.tanh(X)**2) key, subkey1, subkey2 = jax.random.split(key, 3) X = jax.random.normal(subkey1, (30, 40)) V = jax.random.normal(subkey2, (30, 40)) def hessian(f): return jax.jacfwd(jax.jacrev(f)) ans1 = hvp(f, (X,), (V,)) ans2 = jnp.tensordot(hessian(f)(X), V, 2) print(jnp.allclose(ans1, ans2, 1e-4, 1e-4))

如果输出True,说明 forward-over-reverse 的hvp与显式物化 Hessian 再与 $V$ 做二阶张量缩并的结果一致。这里的hessian = jacfwd(jacrev(f))也正是 JAX 内置jax.hessian的实现方式——见 jax/_src/api.py,其源码就是return jacfwd(jacrev(fun, ...), ...),即前向套反向(forward-over-reverse)

4.4 反向套前向与反向套反向

除了 forward-over-reverse,还有另外两种组合方式:

# Reverse-over-forward def hvp_revfwd(f, primals, tangents): g = lambda primals: jax.jvp(f, primals, tangents)[1] return jax.grad(g)(primals)
# Reverse-over-reverse,仅适用于单一参数 def hvp_revrev(f, primals, tangents): x, = primals v, = tangents return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)

官方指南给出的结论是:reverse-over-forward 不如 forward-over-reverse。原因在于:前向模式的开销小于反向模式;而此处外层微分算子要微分的计算比内层更大,因此把开销更小的前向模式放在外层效果最好。

4.5 三种 HVP 的性能对比

下面用 IPython 的%timeit魔法对三种 HVP 及"朴素完整 Hessian 物化"做基准对比(%timeit为 IPython/Jupyter 内置魔法,脚本运行请改用timeit模块):

print("Forward over reverse") %timeit -n10 -r3 hvp(f, (X,), (V,)) print("Reverse over forward") %timeit -n10 -r3 hvp_revfwd(f, (X,), (V,)) print("Reverse over reverse") %timeit -n10 -r3 hvp_revrev(f, (X,), (V,)) print("Naive full Hessian materialization") %timeit -n10 -r3 jnp.tensordot(jax.hessian(f)(X), V, 2)

一般规律是:hvp(forward-over-reverse)最快,hvp_revfwd次之,hvp_revrev再次,而显式物化完整 Hessian 的做法最慢——因为它把 $30 \times 40 \times 30 \times 40$ 的完整 Hessian 张量都算了出来,而 HVP 根本不需要物化它。这一对比也解释了为何在优化算法(如共轭梯度、L-BFGS 的隐式求解)中,HVP 是标准的高效原语。

五、组合 VJP、JVP 与jax.vmap

jax.jvpjax.vjp每次只前推/拉回单个向量。要同时前推/拉回一整批向量,可以借助 JAX 的jax.vmap变换(自动向量化,参见 docs/automatic-vectorization.md),用它写出快速的矩阵-Jacobian 与 Jacobian-矩阵乘积。

5.1 矩阵-Jacobian 乘积(Matrix-Jacobian Product, MJP)

需求:把矩阵 $M$ 的每一行 $m_i$ 作为余切向量,沿 $f$ 在 $W$ 处拉回。先看朴素的 Python 循环版本:

# 隔离出"从权重矩阵到预测值"的函数 f = lambda W: predict(W, b, inputs) # 沿 f 在 W 处拉回余向量 m_i,对 M 的所有行 i # 先用列表推导式在矩阵 M 的行上循环 def loop_mjp(f, x, M): y, vjp_fun = jax.vjp(f, x) return jnp.vstack([jnp.asarray(vjp_fun(mi)) for mi in M]) # 再用 vmap 构造单次快速的矩阵-矩阵乘法, # 而不是外层循环若干次向量-矩阵乘法 def vmap_mjp(f, x, M): y, vjp_fun = jax.vjp(f, x) outs, = jax.vmap(vjp_fun)(M) return outs key = jax.random.key(0) num_covecs = 128 U = jax.random.normal(key, (num_covecs,) + y.shape) loop_vs = loop_mjp(f, W, M=U) print('Non-vmapped Matrix-Jacobian product') %timeit -n10 -r3 loop_mjp(f, W, M=U) print('\nVmapped Matrix-Jacobian product') vmap_vs = vmap_mjp(f, W, M=U) %timeit -n10 -r3 vmap_mjp(f, W, M=U) assert jnp.allclose(loop_vs, vmap_vs), 'Vmap and non-vmapped Matrix-Jacobian Products should be identical'

两者的结果必须一致(代码末尾的assert会校验),但vmap_mjp内部把 128 次独立的向量-矩阵乘法合并成一次矩阵-矩阵乘法,通常能获得数量级的加速。

5.2 Jacobian-矩阵乘积(Jacobian-Matrix Product, JMP)

对称地,可以把矩阵 $M$ 的每一行作为切向量前推:

def loop_jmp(f, W, M): # jvp 会立即以元组形式返回 primal 与 tangent 值, # 因此在列表推导式中计算并选取 tangent 部分 return jnp.vstack([jax.jvp(f, (W,), (mi,))[1] for mi in M]) def vmap_jmp(f, W, M): _jvp = lambda s: jax.jvp(f, (W,), (s,))[1] return jax.vmap(_jvp)(M) num_vecs = 128 S = jax.random.normal(key, (num_vecs,) + W.shape) loop_vs = loop_jmp(f, W, M=S) print('Non-vmapped Jacobian-Matrix product') %timeit -n10 -r3 loop_jmp(f, W, M=S) vmap_vs = vmap_jmp(f, W, M=S) print('\nVmapped Jacobian-Matrix product') %timeit -n10 -r3 vmap_jmp(f, W, M=S) assert jnp.allclose(loop_vs, vmap_vs), 'Vmap and non-vmapped Jacobian-Matrix products should be identical'

注意jax.jvp的返回值是(primals_out, tangents_out)对,所以循环版本里要用[1]取 tangent 部分;vmap_jmp则把"对每行 $s$ 做 JVP 再取 tangent"封装成_jvp,交给jax.vmap批量执行。

六、jax.jacfwdjax.jacrev的实现原理

有了快速 Jacobian-矩阵与矩阵-Jacobian 乘积,jax.jacfwdjax.jacrev的实现思路就呼之欲出了:用同样的技巧,一次性前推或拉回整个标准基(与单位矩阵同构)。

6.1 用vjp+vmap实现jacrev

from jax import jacrev as builtin_jacrev def our_jacrev(f): def jacfun(x): y, vjp_fun = jax.vjp(f, x) # 用 vmap 做矩阵-Jacobian 乘积。 # 这里的矩阵是欧氏基,因此一次得到 Jacobian 的全部元素。 J, = jax.vmap(vjp_fun, in_axes=0)(jnp.eye(len(y))) return J return jacfun assert jnp.allclose(builtin_jacrev(f)(W), our_jacrev(f)(W)), 'Incorrect reverse-mode Jacobian results!'

思路是:vjp先拉回单个余切向量,得到 Jacobian 的一行;把单位矩阵的每一行作为余切批量喂给vjp_fun,就同时得到所有行,即完整 Jacobian。这正对应 jax/_src/api.py 中jacrev的真实实现模式:y, pullback, ... = vjp(f_partial, *dyn_args)之后再jac = vmap(pullback)(_std_basis(y))

6.2 用jvp+vmap实现jacfwd

from jax import jacfwd as builtin_jacfwd def our_jacfwd(f): def jacfun(x): _jvp = lambda s: jax.jvp(f, (x,), (s,))[1] Jt = jax.vmap(_jvp, in_axes=1)(jnp.eye(len(x))) return jnp.transpose(Jt) return jacfun assert jnp.allclose(builtin_jacfwd(f)(W), our_jacfwd(f)(W)), 'Incorrect forward-mode Jacobian results!'

jax.jvp一次前推单个切向量,得到 Jacobian 的一列。为了高效地把单位矩阵的所有列一起前推,这里用in_axes=1jnp.eye(len(x))映射为vmap的批量维度;由于vmap输出的每一"行"对应一列 Jacobian(即转置后的结果),最后用jnp.transpose还原。JAX 内置jacfwd的实现与之同理(jax/_src/api.py),它通过vmap(pushfwd, out_axes=(None, -1))(_std_basis(dyn_args))一次推过整个标准基。

补充:源码中构造标准基的工具是_std_basis(jax/_src/api.py),它把 pytree 展平后调用jnp.eye(ndim)生成单位矩阵作为基。另外,jax.jacobian只是jax.jacrev的别名(jax/_src/api.py),需要前向模式请显式使用jax.jacfwd

6.3 为什么 Autograd 做不到这些

有趣的是,Autograd 库做不到上述实现。Autograd 的反向模式jacobian只能通过外层循环的map逐次拉回单个向量;一次只把一个向量推过整个计算,远比用jax.vmap把整个批次合并起来计算低效。这正是jax.vmap与 JAX 的jaxpr追踪机制带来的优势。

6.4 微分计算的线性部分可以被 JIT

Autograd 做不到的另一件事是jax.jit。有趣的是,无论被微分的函数里使用了多少 Python 动态行为,我们总能对计算的线性部分使用jax.jit。例如:

def f(x): try: if x < 3: return 2 * x ** 3 else: raise ValueError except ValueError: return jnp.pi * x y, f_vjp = jax.vjp(f, 4.) print(jax.jit(f_vjp)(1.))

f内部有try/except和条件分支,追踪阶段就已经把 Python 控制流"消化"掉了;f_vjp是纯线性函数,因此可以安全地交给jax.jit编译执行。对vjp函数做 JIT 的机制在 JAX 的测试套件(如 tests/jax_jit_test.py、tests/lax_autodiff_test.py)中都有覆盖,可用作进一步验证。

七、小结与延伸阅读

两种模式的核心取舍可以浓缩为一张对照表:

维度前向模式(JVP,jax.jvp反向模式(VJP,jax.vjp
基本运算$(x, v) \mapsto (f(x), \partial f(x) v)$$(x, v) \mapsto (f(x), v^\mathsf{T}\partial f(x))$
构建完整 Jacobian逐列(one-hot 切向量)逐行(one-hot 余切向量)
FLOPs约 3 倍于函数求值约 3 倍于函数求值
内存与计算深度无关随计算深度增长
典型场景输出维度远小于输入维度("高瘦" Jacobian)、HVP标量损失梯度("宽扁" Jacobian)、jax.grad
实现复杂度相对简单(边算边推)相对复杂(需保存中间值做反向传播)

jax.jacfwd/jax.jacrev/jax.hessian都是上述原语的组合:hessianjacfwd(jacrev(f))(前向套反向);自定义 Hessian-vector product 的最佳实践是jax.jvp(jax.grad(f), primals, tangents)[1];需要批量前推/拉回时,用jax.vmapjvp/vjp组合成矩阵-Jacobian / Jacobian-矩阵乘积。

如果想继续深入,仓库内还有大量相关资源:更进阶的自动微分讨论见 docs/advanced_autodiff.md;完整的 autodiff 实践手册见 docs/autodiff_cookbook.md;从零实现 JAX 风格自动微分的教程见 docs/autodidax.md;底层 AD 解释器源码在 jax/_src/interpreters/ad.py,jax.grad的 docstring 也明确指出它是jax.vjp的特例(jax/_src/api.py)。

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询