资讯动态

JAX 自动微分 Cookbook:从 grad 到 JVP/VJP 的前向与反向模式全解

发布时间:2026/9/10 7:46:01 来源:尧图企业网站定制
JAX 自动微分 Cookbook从 grad 到 JVP/VJP 的前向与反向模式全解【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本篇技术指南以 JAX 官方《The Autodiff Cookbook》为骨架系统讲解 JAX 自动微分系统的核心用法与底层原理从最基础的grad、value_and_grad到高阶导数Hessian-vector product、jacfwd/jacrev与稠密 Hessian再到支撑这一切的两个奠基性原语——前向模式的 Jacobian-vector productJVP与反向模式的 vector-Jacobian productVJP最后深入复数域的微分规则。读完本文你将掌握在 JAX 中构建任意阶、任意模式微分程序的完整工具箱并理解每条 API 在仓库源码中的真实实现路径。本文所有代码与结论均以当前仓库为准关键 API 的实现可直接在 jax/_src/api.py 中核对配套的可执行 Notebook 版本见 docs/notebooks/autodiff_cookbook.ipynb。准备工作本文示例使用如下导入注意当前仓库使用类型化随机密钥typed keys因此以random.key(0)而非旧式random.PRNGKey(0)构造密钥import jax.numpy as jnp from jax import grad, jit, vmap from jax import random key random.key(0)从grad开始梯度计算基础grad的基本用法可以用grad对任意函数求导grad_tanh grad(jnp.tanh) print(grad_tanh(2.0))grad的输入输出都是函数如果 Python 函数f计算数学函数 $f$那么grad(f)就是一个计算 $\nabla f$ 的 Python 函数即grad(f)(x)的值为 $\nabla f(x)$。由于grad作用于函数可以把它反复套在自己的输出上求任意阶导数print(grad(grad(jnp.tanh))(2.0)) print(grad(grad(grad(jnp.tanh)))(2.0))在仓库源码中grad是value_and_grad的薄封装见 jax/_src/api.pygrad先调用value_and_grad构造同时返回函数值与梯度的包装函数再丢弃函数值、只返回梯度has_auxTrue时则返回(gradient, aux)二元组。实战示例线性逻辑回归的梯度下面用一个完整的线性逻辑回归模型演示梯度计算。首先定义模型与损失def sigmoid(x): return 0.5 * (jnp.tanh(x / 2) 1) # 输出标签为 True 的概率 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]]) targets jnp.array([True, True, False, True]) # 训练损失为训练样本的负对数似然 def loss(W, b): preds predict(W, b, inputs) label_probs preds * targets (1 - preds) * (1 - targets) return -jnp.sum(jnp.log(label_probs)) # 初始化随机模型系数 key, W_key, b_key random.split(key, 3) W random.normal(W_key, (3,)) b random.normal(b_key, ())argnums指定对哪些位置参数求导grad的argnums参数用来指定对函数的位置参数positional arguments中的哪一个求导# 对第一个位置参数求导 W_grad grad(loss, argnums0)(W, b) print(W_grad, W_grad) # 因为 argnums0 是默认值以下写法等价 W_grad grad(loss)(W, b) print(W_grad, W_grad) # 也可以选择其他参数并省略关键字写法 b_grad grad(loss, 1)(W, b) print(b_grad, b_grad) # 支持元组形式同时求多个参数的梯度 W_grad, b_grad grad(loss, (0, 1))(W, b) print(W_grad, W_grad) print(b_grad, b_grad)这个gradAPI 与 Spivak 的经典教材Calculus on Manifolds1965以及 Sussman 与 Wisdom 的Structure and Interpretation of Classical Mechanics2015、Functional Differential Geometry2013中的记号一一对应。本质上argnums使grad(f, i)成为计算偏导数 $\partial_i f$ 的函数。从源码看argnums可以是一个整数或整数序列jax/_src/api.py 中通过argnums_partial2将目标参数从args/kwargs中提取出来作为动态参数参与微分其余参数被当作常量冻结如果传入的argnums超出实际位置参数个数会直接抛出TypeError。此外grad还支持以下参数均定义于 jax/_src/api.pyhas_aux若fun返回(主输出, 辅助数据)二元组设为True后grad返回(梯度, 辅助数据)holomorphic承诺被微分的函数是复解析全纯函数要求输入输出均为复数 dtypeallow_int允许对整数输入求导此时梯度为平凡的float0向量空间 dtype默认False。注意grad只接受输出为标量包括 shape 为()的数组但不包括 shape(1,)等的函数输出非标量时会得到带提示的TypeError建议改用jax.jacobian见 jax/_src/api.py 中的_check_scalar。对嵌套 list、tuple 和 dict 求导对标准 Python 容器求导开箱即用tuple、list、dict 以及任意嵌套都可以直接使用def loss2(params_dict): preds predict(params_dict[W], params_dict[b], inputs) label_probs preds * targets (1 - preds) * (1 - targets) return -jnp.sum(jnp.log(label_probs)) print(grad(loss2)({W: W, b: b}))这是 JAX pytree 机制的功劳任意容器被展平为叶子数组的树梯度保持相同的树结构返回。你还可以注册自定义容器类型使其不仅适用于grad也适用于所有 JAX 变换jit、vmap等。用value_and_grad同时求函数值与梯度value_and_grad可以一次调用同时高效地返回函数值与梯度from jax import value_and_grad loss_value, Wb_grad value_and_grad(loss, (0, 1))(W, b) print(loss value, loss_value) print(loss value, loss(W, b))从源码看value_and_grad是比grad更底层的实现它先对部分应用后的函数调用vjp得到(ans, vjp_py)再把单位余切向量lax_internal._one_vjp(ans)送入 VJP 得到梯度见 jax/_src/api.py。这正是grad是vjp的一个特例这一事实的实现体现。用数值差分校验梯度导数的一个好处是可以用有限差分直接验证。对上述模型# 有限差分步长 eps 1e-4 # 对标量 b_grad 做标量有限差分 b_grad_numerical (loss(W, b eps / 2.) - loss(W, b - eps / 2.)) / eps print(b_grad_numerical, b_grad_numerical) print(b_grad_autodiff, grad(loss, 1)(W, b)) # 对 W_grad 沿随机方向做有限差分 key, subkey random.split(key) vec random.normal(subkey, W.shape) unitvec vec / jnp.sqrt(jnp.vdot(vec, vec)) W_grad_numerical (loss(W eps / 2. * unitvec, b) - loss(W - eps / 2. * unitvec, b)) / eps print(W_dirderiv_numerical, W_grad_numerical) print(W_dirderiv_autodiff, jnp.vdot(grad(loss)(W, b), unitvec))JAX 还提供了封装好的便捷函数可以校验到任意阶导数from jax.test_util import check_grads check_grads(loss, (W, b), order2) # 校验到 2 阶导数check_grads的实现位于 jax/_src/public_test_util.py它支持modes(fwd, rev, lin)三种校验模式默认同时校验前向JVP与反向VJP递归地对order阶以内的导数逐一做 JVP/VJP 与数值差分的一致性检查梯度只在一个随机选择的方向上校验从而避免大输入/输出空间下有限差分代价过高。用grad-of-grad构造 Hessian-vector product利用高阶grad可以构造 Hessian-vector productHVP函数。HVP 在截断牛顿共轭梯度算法truncated Newton conjugate-gradient中用于最小化光滑凸函数也常用于研究神经网络训练目标的曲率。对于具有连续二阶导数的标量函数 $f : \mathbb{R}^n \to \mathbb{R}$Hessian 矩阵对称Hessian-vector product 函数需要能够对任意 $v \in \mathbb{R}^n$ 计算$\qquad v \mapsto \partial^2 f(x) \cdot v$关键在于不要物化完整的 Hessian 矩阵当 $n$ 达到数百万甚至数十亿神经网络场景时存储完整 Hessian 是不可能的。幸运的是利用恒等式$\qquad \partial^2 f (x) v \partial [x \mapsto \partial f(x) \cdot v] \partial g(x)$其中 $g(x) \partial f(x) \cdot v$ 是把 $f$ 在 $x$ 处的梯度与向量 $v$ 做点积得到的新标量函数。注意我们全程只在对向量值参数上的标量函数求导这正是grad最擅长的场景。JAX 代码如下def hvp(f, x, v): return grad(lambda x: jnp.vdot(grad(f)(x), v))(x)这个例子也说明 JAX 对词法闭包lexical closure的处理非常稳健任意嵌套都不会混淆。用jacfwd与jacrev计算 Jacobian 与 Hessian计算稠密 Jacobianjacfwd和jacrev可以计算完整的 Jacobian 矩阵from jax import jacfwd, jacrev # 隔离出从权重矩阵到预测的函数 f lambda W: predict(W, b, inputs) J jacfwd(f)(W) print(jacfwd result, with shape, J.shape) print(J) J jacrev(f)(W) print(jacrev result, with shape, J.shape) print(J)两个函数计算相同的数值结果差异只在机器数值精度范围内但实现方式不同jacfwd使用前向模式自动微分对高瘦型 Jacobian输出数多于输入数更高效jacrev使用反向模式对宽扁型 Jacobian输入数多于输出数更高效接近方阵时jacfwd通常略占优势。jacfwd/jacrev同样支持容器类型def predict_dict(params, inputs): return predict(params[W], params[b], inputs) J_dict jacrev(predict_dict)({W: W, b: b}, inputs) for k, v in J_dict.items(): print(Jacobian from {} to logits is.format(k)) print(v)用jacfwd(jacrev(f))计算稠密 Hessian组合两个函数即可计算稠密 Hessian 矩阵def hessian(f): return jacfwd(jacrev(f)) H hessian(f)(W) print(hessian, with shape, H.shape) print(H)形状的直觉对 $f : \mathbb{R}^n \to \mathbb{R}^m$在点 $x$ 处依次有$f(x) \in \mathbb{R}^m$$f$ 在 $x$ 处的值$\partial f(x) \in \mathbb{R}^{m \times n}$$x$ 处的 Jacobian 矩阵$\partial^2 f(x) \in \mathbb{R}^{m \times n \times n}$$x$ 处的 Hessian以此类推。实现hessian时理论上可以用jacfwd(jacrev(f))、jacrev(jacfwd(f))或任意组合但前向套反向forward-over-reverse通常最高效内层 Jacobian 常常是在对宽扁函数求导如损失函数 $f : \mathbb{R}^n \to \mathbb{R}$反向模式占优而外层 Jacobian 是在对 $\nabla f : \mathbb{R}^n \to \mathbb{R}^n$ 求导其 Jacobian 是方阵此时前向模式胜出。仓库中jax.hessian正是以此为默认实现见 jax/_src/api.py。基石之一Jacobian-vector productJVP前向模式JAX 同时实现了高效且通用的前向与反向自动微分。grad建立在反向模式之上但要理解两种模式的差异与各自适用场景需要一点数学背景。JVP 的数学定义给定函数 $f : \mathbb{R}^n \to \mathbb{R}^m$$f$ 在输入点 $x$ 处的 Jacobian $\partial f(x)$ 通常被视为 $\mathbb{R}^{m \times n}$ 矩阵但也可以看作一个线性映射——把 $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 map推前映射Jacobian 矩阵就是该线性映射在标准基下的矩阵表示。若不对具体输入点 $x$ 做承诺可以把 $\partial f$ 视为先取输入点、再返回该点的 Jacobian 线性映射的函数。将输入点 $x \in \mathbb{R}^n$ 与切向量 $v \in \mathbb{R}^n$ 配对得到输出切向量在 $\mathbb{R}^m$ 中这个从 $(x, v)$ 到输出切向量的映射就是Jacobian-vector product记作$\qquad (x, v) \mapsto \partial f(x) v$JVP 的 JAX 代码JAX 的jvp函数正是这一变换的建模给定计算 $f$ 的 Python 函数jvp给出计算 $(x, v) \mapsto (f(x), \partial f(x) v)$ 的 Python 函数。from jax import jvp # 隔离出从权重矩阵到预测的函数 f lambda W: predict(W, b, inputs) key, subkey random.split(key) v random.normal(subkey, W.shape) # 将向量 v 沿 f 在 W 处推前 y, u jvp(f, (W,), (v,))用类 Haskell 的类型签名可以写作jvp :: (a - b) - a - T a - (b, T b)其中T a表示a的切空间类型。即jvp接收一个a - b函数、一个a类型的值、一个T a类型的切向量返回由b类型值与T b类型输出切向量组成的二元组。jvp变换后的函数求值方式与原函数类似只是每个 primal 值旁边都伴随一个T a类型的切向量。对于原函数执行的每个原始数值运算jvp变换后的函数都会执行该原语的JVP 规则既在 primal 上求值原语又在这些 primal 值上应用该原语的 JVP。这种求值策略直接决定了计算复杂度由于 JVP 是边算边走无需为后续步骤保存任何中间量内存开销与计算的深度无关同时 JVP 变换后函数的 FLOP 开销约为原函数的 3 倍一份用于求值原函数如sin(x)一份用于线性化如cos(x)一份用于把线性化结果作用到向量上如cos_x * v。换言之固定 primal 点 $x$ 后计算 $v \mapsto \partial f(x) \cdot v$ 的边际成本与求值 $f$ 相当。为什么机器学习中不常用前向模式既然内存复杂度如此诱人为什么前向模式在机器学习中很少见思考如何用 JVP 构建完整 Jacobian把 JVP 作用到 one-hot 切向量上就会揭示 Jacobian 中与该非零分量对应的一列。于是可以一列一列地构建完整 Jacobian每列的成本约等于一次函数求值。这对高瘦 Jacobian 高效对宽扁 Jacobian 则低效。而在基于梯度的机器学习优化中目标通常是从 $\mathbb{R}^n$ 的参数空间到标量损失 $\mathbb{R}$其 Jacobian 是极宽的矩阵 $\partial f(x) \in \mathbb{R}^{1 \times n}$通常被等同于梯度向量 $\nabla f(x) \in \mathbb{R}^n$。一列一列地构建该矩阵、每次调用消耗与原始函数相近的 FLOP显然低效。尤其对于神经网络训练损失$n$ 可达数百万甚至数十亿这种方式根本无法扩展。要做得更好就需要反向模式。基石之二Vector-Jacobian productVJP反向模式前向模式返回计算 Jacobian-vector product 的函数可以一列一列构建 Jacobian反向模式则返回计算 vector-Jacobian product等价于 Jacobian 转置乘向量的函数可以一行一行构建 Jacobian。VJP 的数学定义再次考虑 $f : \mathbb{R}^n \to \mathbb{R}^m$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)$但常见情形下可以把 $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$ 的输入正如对转置线性函数的预期。VJP 的 JAX 代码JAX 函数vjp接收计算 $f$ 的 Python 函数返回计算 VJP $(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 random.split(key) u random.normal(subkey, y.shape) # 将余向量 u 沿 f 在 W 处拉回 v vjp_fun(u)类 Haskell 类型签名vjp :: (a - b) - a - (b, CT b - CT a)其中CT a表示a的余切空间类型。即vjp接收a - b函数与a类型点返回由b类型值与CT b - CT a线性映射组成的二元组。这正是它的价值可以一行一行构建 Jacobian且计算 $(x, v) \mapsto (f(x), v^\mathsf{T} \partial f(x))$ 的 FLOP 开销仅为求值 $f$ 的约 3 倍。特别地对 $f : \mathbb{R}^n \to \mathbb{R}$ 求梯度只需一次调用——这就是grad对基于梯度的优化如此高效的原因哪怕目标函数是拥有数百万乃至数十亿参数的神经网络训练损失。代价是虽然 FLOP 友好内存随计算的深度增长且实现传统上比前向模式更复杂。在仓库源码中vjp的完整实现见 jax/_src/api.py其返回值vjpfun接收与primals_out形状相同的余切向量返回与primals数量、形状相同的余切向量元组它还支持has_aux、saveable_args、in_nzs等进阶参数。用 VJP 实现向量值梯度如果你需要向量值梯度类似tf.gradients的语义from jax import vjp def vgrad(f, x): y, vjp_fn vjp(f, x) return vjp_fn(jnp.ones(y.shape))[0] print(vgrad(lambda x: 3*x**2, jnp.ones((2, 2))))前向-反向混合的 Hessian-vector product之前我们用纯反向模式实现了 HVP假设二阶导连续def hvp(f, x, v): return grad(lambda x: jnp.vdot(grad(f)(x), v))(x)它已经高效但把前向模式与反向模式结合起来还能进一步节省内存。数学上对 $f : \mathbb{R}^n \to \mathbb{R}$、线性化点 $x$ 与向量 $v$我们想要$(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$几乎可以直接翻译成代码from jax import jvp, grad # forward-over-reverse前向套反向 def hvp(f, primals, tangents): return jvp(grad(f), primals, tangents)[1]更好的地方在于由于不需要直接调用jnp.dot这个hvp对任意形状的数组、任意容器类型如存为嵌套 list/dict/tuple 的向量都适用甚至不依赖jax.numpy。使用示例def f(X): return jnp.sum(jnp.tanh(X)**2) key, subkey1, subkey2 random.split(key, 3) X random.normal(subkey1, (30, 40)) V random.normal(subkey2, (30, 40)) ans1 hvp(f, (X,), (V,)) ans2 jnp.tensordot(hessian(f)(X), V, 2) print(jnp.allclose(ans1, ans2, 1e-4, 1e-4))另一种写法是反向套前向reverse-over-forward# reverse-over-forward def hvp_revfwd(f, primals, tangents): g lambda primals: jvp(f, primals, tangents)[1] return grad(g)(primals)它稍逊一筹因为前向模式的开销小于反向模式而外层的微分算子需要微分的计算规模比内层更大所以把前向模式放在外层效果最佳# reverse-over-reverse仅适用于单参数 def hvp_revrev(f, primals, tangents): x, primals v, tangents return grad(lambda x: jnp.vdot(grad(f)(x), v))(x) 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(hessian(f)(X), V, 2)组合 VJP、JVP 与vmapJacobian-Matrix 与 Matrix-Jacobian 乘积有了jvp与vjp这两个单次推前/拉回一个向量的变换就可以用vmap一次性推前/拉回整个基。先看 Matrix-Jacobian 乘积拉回多个余向量# 隔离出从权重矩阵到预测的函数 f lambda W: predict(W, b, inputs) # 沿 f 在 W 处拉回余向量 m_i对矩阵 M 的所有行 i。 # 先用列表推导式在矩阵 M 的行上循环。 def loop_mjp(f, x, M): y, vjp_fun 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 vjp(f, x) outs, vmap(vjp_fun)(M) return outs key random.key(0) num_covecs 128 U random.normal(key, (num_covecs,) y.shape) loop_vs loop_mjp(f, W, MU) print(Non-vmapped Matrix-Jacobian product) %timeit -n10 -r3 loop_mjp(f, W, MU) print(\nVmapped Matrix-Jacobian product) vmap_vs vmap_mjp(f, W, MU) %timeit -n10 -r3 vmap_mjp(f, W, MU) assert jnp.allclose(loop_vs, vmap_vs), Vmap and non-vmapped Matrix-Jacobian Products should be identical再看 Jacobian-Matrix 乘积推前多个切向量def loop_jmp(f, W, M): # jvp 立即返回 primal 与 tangent 组成的元组 # 因此我们在列表推导式中计算并选取 tangent 部分 return jnp.vstack([jvp(f, (W,), (mi,))[1] for mi in M]) def vmap_jmp(f, W, M): _jvp lambda s: jvp(f, (W,), (s,))[1] return vmap(_jvp)(M) num_vecs 128 S random.normal(key, (num_vecs,) W.shape) loop_vs loop_jmp(f, W, MS) print(Non-vmapped Jacobian-Matrix product) %timeit -n10 -r3 loop_jmp(f, W, MS) vmap_vs vmap_jmp(f, W, MS) print(\nVmapped Jacobian-Matrix product) %timeit -n10 -r3 vmap_jmp(f, W, MS) assert jnp.allclose(loop_vs, vmap_vs), Vmap and non-vmapped Jacobian-Matrix products should be identicaljacfwd与jacrev的实现有了上述快速的 Jacobian-Matrix 与 Matrix-Jacobian 乘积jacfwd/jacrev的实现思路就非常直接了用同样的技术一次性推前/拉回一整个标准基与单位矩阵同构。先看反向模式版from jax import jacrev as builtin_jacrev def our_jacrev(f): def jacfun(x): y, vjp_fun vjp(f, x) # 用 vmap 做 Matrix-Jacobian 乘积。 # 这里矩阵是欧几里得基因此一次得到 Jacobian 的所有元素。 J, vmap(vjp_fun, in_axes0)(jnp.eye(len(y))) return J return jacfun assert jnp.allclose(builtin_jacrev(f)(W), our_jacrev(f)(W)), Incorrect reverse-mode Jacobian results!再看前向模式版from jax import jacfwd as builtin_jacfwd def our_jacfwd(f): def jacfun(x): _jvp lambda s: jvp(f, (x,), (s,))[1] Jt vmap(_jvp, in_axes1)(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!这与仓库源码的实现思路完全一致jacfwd本质上是vmap(_jvp)(_std_basis(...))加上_jacfwd_unravel的重组见 jax/_src/api.pyjacrev则是vjp之后用vmap(pullback)(_std_basis(y))一次拉回整个标准基见 jax/_src/api.py。值得一提的还有jax.hessian正是jacfwd(jacrev(fun))的别名jax/_src/api.py其默认采用 forward-over-reverse 组合。一个有趣的历史背景是早期的 Autograd 库做不到这一点——它实现反向模式jacobian时必须用外层map循环一次只拉回一个向量。一次一个向量地穿过计算图远不如用vmap把所有向量批量合并起来高效。Autograd 做不到的另一件事是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 vjp(f, 4.) print(jit(f_vjp)(1.))复数与微分JAX 对复数和微分支持良好。为同时支持全纯holomorphic与非全纯微分最好以 JVP 和 VJP 的视角来思考。考虑复到复函数 $f: \mathbb{C} \to \mathbb{C}$并把它等同于相应的 $g: \mathbb{R}^2 \to \mathbb{R}^2$def f(z): x, y jnp.real(z), jnp.imag(z) return u(x, y) v(x, y) * 1j def g(x, y): return (u(x, y), v(x, y))即分解 $f(z) u(x, y) v(x, y) i$其中 $z x y i$把 $\mathbb{C}$ 等同于 $\mathbb{R}^2$ 得到 $g$。由于 $g$ 只涉及实数输入输出我们已经知道如何为它写 Jacobian-vector product给定切向量 $(c, d) \in \mathbb{R}^2$即$\begin{bmatrix} \partial_0 u(x, y) \partial_1 u(x, y) \ \partial_0 v(x, y) \partial_1 v(x, y) \end{bmatrix} \begin{bmatrix} c \ d \end{bmatrix}$要对原函数 $f$ 作用于切向量 $c di \in \mathbb{C}$ 得到 JVP只需沿用同一套定义并把结果等同为另一个复数$\partial f(x y i)(c d i) \begin{bmatrix} 1 i \end{bmatrix} \begin{bmatrix} \partial_0 u(x, y) \partial_1 u(x, y) \ \partial_0 v(x, y) \partial_1 v(x, y) \end{bmatrix} \begin{bmatrix} c \ d \end{bmatrix}$这就是 $\mathbb{C} \to \mathbb{C}$ 函数 JVP 的定义注意 $f$ 是否全纯无关紧要JVP 是无歧义的。下面做一次验证def check(seed): key random.key(seed) # 为 u 和 v 生成随机系数 key, subkey random.split(key) a, b, c, d random.uniform(subkey, (4,)) def fun(z): x, y jnp.real(z), jnp.imag(z) return u(x, y) v(x, y) * 1j def u(x, y): return a * x b * y def v(x, y): return c * x d * y # primal 点 key, subkey random.split(key) x, y random.uniform(subkey, (2,)) z x y * 1j # 切向量 key, subkey random.split(key) c, d random.uniform(subkey, (2,)) z_dot c d * 1j # 检查 jvp _, ans jvp(fun, (z,), (z_dot,)) expected (grad(u, 0)(x, y) * c grad(u, 1)(x, y) * d grad(v, 0)(x, y) * c * 1j grad(v, 1)(x, y) * d * 1j) print(jnp.allclose(ans, expected))check(0) check(1) check(2)那么 VJP 呢做法类似对余切向量 $c di \in \mathbb{C}$把 $f$ 的 VJP 定义为$(c di)^* ; \partial f(x y i) \begin{bmatrix} c -d \end{bmatrix} \begin{bmatrix} \partial_0 u(x, y) \partial_1 u(x, y) \ \partial_0 v(x, y) \partial_1 v(x, y) \end{bmatrix} \begin{bmatrix} 1 \ -i \end{bmatrix}$为什么有负号只是为了处理复共轭以及我们是在与余向量covector打交道。VJP 规则的验证def check(seed): key random.key(seed) # 为 u 和 v 生成随机系数 key, subkey random.split(key) a, b, c, d random.uniform(subkey, (4,)) def fun(z): x, y jnp.real(z), jnp.imag(z) return u(x, y) v(x, y) * 1j def u(x, y): return a * x b * y def v(x, y): return c * x d * y # primal 点 key, subkey random.split(key) x, y random.uniform(subkey, (2,)) z x y * 1j # 余切向量 key, subkey random.split(key) c, d random.uniform(subkey, (2,)) z_bar jnp.array(c d * 1j) # 用于控制 dtype # 检查 vjp _, fun_vjp vjp(fun, z) ans, fun_vjp(z_bar) expected (grad(u, 0)(x, y) * c grad(v, 0)(x, y) * (-d) grad(u, 1)(x, y) * c * (-1j) grad(v, 1)(x, y) * (-d) * (-1j)) assert jnp.allclose(ans, expected, atol1e-5, rtol1e-5)check(0) check(1) check(2)grad、jacfwd、jacrev的复数行为回忆对 $\mathbb{R} \to \mathbb{R}$ 函数grad(f)(x)被定义为vjp(f, x)1——把 VJP 作用到值1.0上即可揭示梯度即 Jacobian即导数。对 $\mathbb{C} \to \mathbb{R}$ 函数可以做同样的事仍然使用1.0作为余切向量得到的复数结果概括了完整 Jacobiandef f(z): x, y jnp.real(z), jnp.imag(z) return x**2 y**2 z 3. 4j grad(f)(z)对一般的 $\mathbb{C} \to \mathbb{C}$ 函数Jacobian 有 4 个实数自由度如上述 2×2 Jacobian 矩阵无法全部塞进一个复数。但对全纯函数可以全纯函数正是导数能用单个复数表示的 $\mathbb{C} \to \mathbb{C}$ 函数Cauchy-Riemann 方程保证上述 2×2 Jacobian 具有复平面中缩放-旋转矩阵的特殊形式即单个复数乘法的作用。而这个复数可以用一次vjp余向量取1.0揭示出来。由于这一技巧只对全纯函数成立使用前需要向 JAX 承诺函数是全纯的否则 JAX 会在对复数输出函数使用grad时报错def f(z): return jnp.sin(z) z 3. 4j grad(f, holomorphicTrue)(z)holomorphicTrue承诺的全部作用就是关闭输出为复数时的错误。也可以对并非全纯的函数写holomorphicTrue但得到的答案并不代表完整 Jacobian而是丢弃输出虚部后的函数的 Jacobiandef f(z): return jnp.conjugate(z) z 3. 4j grad(f, holomorphicTrue)(z) # f 实际上并不是全纯的由此grad在复数域有几个实用的推论可以对全纯的 $\mathbb{C} \to \mathbb{C}$ 函数使用grad。可以用grad优化 $f : \mathbb{C} \to \mathbb{R}$ 函数如以复数参数 $x$ 的实值损失函数方法是沿grad(f)(x)的共轭方向步进。如果一个 $\mathbb{R} \to \mathbb{R}$ 函数内部恰好用了某些复数运算其中一些必然非全纯例如卷积中使用的 FFTgrad依然可用且结果与纯实数实现给出的结果一致。无论如何JVP 与 VJP 总是无歧义的若要计算非全纯 $\mathbb{C} \to \mathbb{C}$ 函数的完整 Jacobian 矩阵用 JVP 或 VJP 即可。可以预期复数在 JAX 中处处可用。下面是对复矩阵做 Cholesky 分解并求导的示例A jnp.array([[5., 2.3j, 5j], [2.-3j, 7., 1.7j], [-5j, 1.-7j, 12.]]) def f(X): L jnp.linalg.cholesky(X) return jnp.sum((L - jnp.sin(L))**2) grad(f, holomorphicTrue)(A)进阶自动微分展望本文从易到难、循序渐进地演示了 JAX 中自动微分的各种应用。若想进一步深入仓库内的 docs/advanced_autodiff.md 提供了高级自动微分的完整进阶指南。自动微分的世界里还有大量其他技巧与功能本文未覆盖、未来可能在Advanced Autodiff Cookbook中展开的主题包括Gauss-Newton Vector Products只线性化一次的 Gauss-Newton 向量积Custom VJPs and JVPs自定义 VJP 与 JVP 规则仓库中jax.custom_jvp/jax.custom_vjp的实现位于 jax/_src/custom_derivatives.py对应教程见 docs/notebooks/Custom_derivative_rules_for_Python_code.ipynbEfficient derivatives at fixed-points不动点处的高效求导Estimating the trace of a Hessian用随机 Hessian-vector product 估计 Hessian 的迹Forward-mode autodiff using only reverse-mode只用反向模式实现前向模式自动微分Taking derivatives with respect to custom data types对自定义数据类型求导Checkpointing二项式 checkpointing为高效反向模式而非模型快照仓库中jax.checkpoint/jax.remat的实现在 jax/_src/ad_checkpoint.pyOptimizing VJPs with Jacobian pre-accumulation通过 Jacobian 预累加优化 VJP。关键 API 速查与源码索引API作用源码位置jax.grad对标量输出函数求梯度argnums/has_aux/holomorphic/allow_intjax/_src/api.pyjax.value_and_grad同时返回函数值与梯度jax/_src/api.pyjax.jacfwd前向模式逐列计算 Jacobianjax/_src/api.pyjax.jacrev反向模式逐行计算 Jacobianjax.jacobian是其别名jax/_src/api.pyjax.hessianjacfwd(jacrev(fun))的稠密 Hessian 实现jax/_src/api.pyjax.jvp前向模式 Jacobian-vector productjax/_src/api.pyjax.vjp反向模式 vector-Jacobian productgrad的底层基石jax/_src/api.pyjax.linearize基于jvp与部分求值的线性化工具jax/_src/api.pyjax.test_util.check_grads与数值差分对比校验到任意阶导数jax/_src/public_test_util.py整体而言JAX 的自动微分体系可以用一条简洁的调用链概括grad是value_and_grad的薄封装value_and_grad内部调用vjp并把单位余切向量送入 VJPjacrev用vmap把vjp的 pullback 一次性作用到整个标准基jacfwd用vmap把jvp一次性推前整个标准基hessian则是二者的组合。理解这条调用链与 JVP/VJP 两种模式的数学定义就掌握了在 JAX 中编写任意自动微分程序的核心方法论。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价