ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

JAX 高阶导数完全指南:Hessian、HVP 与 stop_gradient 实战

JAX 高阶导数完全指南:Hessian、HVP 与 stop_gradient 实战 JAX 高阶导数完全指南Hessian、HVP 与 stop_gradient 实战【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读本文基于 JAX 官方教程 docs/higher-order.md 整理而成系统讲解如何在 JAX 中通过叠加变换stacking transformations计算二阶及更高阶导数——包括用jax.jacfwd/jax.jacrev组合实现 Hessian 矩阵、用jax.grad嵌套实现 Hessian-vector product、用jax.lax.stop_gradient截断反向传播实现 TD(0) 强化学习更新与直通估计器以及用jax.jitjax.vmapjax.grad三变换组合高效计算逐样本梯度。读完本文你将掌握 JAX 高阶自动微分从原理到实战的完整方法论并能直接复用文中的代码模式解决元学习、强化学习、二阶优化与逐样本梯度等实际问题。JAX 的自动微分之所以能轻松支撑高阶导数核心在于计算导数的函数本身也是可微分的。因此高阶导数在 JAX 中就是变换套变换无需任何特殊处理。单变量情形在 自动微分教程 中已有覆盖那里用jax.grad计算了 $f(x) x^3 2x^2 - 3x 1$ 的导数本文则重点展开多变量情形下的二阶及高阶导数。Hessian 矩阵Jacobian of the gradient对于多变量情形二阶导数由 Hessian 矩阵 表示其定义为$$(\mathbf{H}f)_{i,j} \frac{\partial^2 f}{\partial_i\partial_j}.$$Hessian 的一个关键事实是实值多元函数 $f: \mathbb{R}^n\to\mathbb{R}$ 的 Hessian 可以等同于其梯度gradient的 Jacobian。于是问题被拆解为两步先求梯度再对梯度求 Jacobian。JAX 为此提供了两个求 Jacobian 的变换jax.jacfwd前向模式自动微分与jax.jacrev反向模式自动微分。它们给出相同的数学结果但在不同场景下效率各异——前向模式在输出多于输入tall Jacobian时更高效反向模式在输入多于输出wide Jacobian时更高效。利用这一性质仅需一行代码即可定义 Hessianimport jax def hessian(f): return jax.jacfwd(jax.grad(f))用点积函数验证 Hessian 正确性以点积函数 $f: \mathbf{x} \mapsto \mathbf{x}^\top \mathbf{x}$ 为例做正确性检验。解析地看当 $ij$ 时 $\frac{\partial^2 f}{\partial_i\partial_j}(\mathbf{x}) 2$否则为 $0$。代码验证如下import jax.numpy as jnp def f(x): return jnp.dot(x, x) hessian(f)(jnp.array([1., 2., 3.]))运行结果应为 $2$ 倍单位矩阵即对角元为 $2$、其余为 $0$ 的 $3\times 3$ 矩阵与解析结果完全一致。源码印证jax.grad在 jax/_src/api.py 中实现它内部委托给value_and_grad支持argnums对第几个位置参数求导可为整数或序列、has_aux函数返回 (输出, 辅助数据) 对、holomorphic复值全纯函数微分、allow_int允许对整数输入求导梯度为 float0 平凡向量空间 dtype等参数。jax.jacfwd在 jax/_src/api.py 中实现其核心思路是对每个标准基向量方向构造 JVP_jvp再用jax.vmap批量推进从而逐列得到 Jacobian。高阶导数的典型应用MAML 元学习一些元学习技术需要在梯度更新之后继续求导。例如 Model-Agnostic Meta-LearningMAML其核心思想是让模型在若干步梯度更新后的损失最小化因此必须对梯度更新本身求导。在其他框架中这往往相当繁琐但在 JAX 中只需自然地编写代码def meta_loss_fn(params, data): Computes the loss after one step of SGD. grads jax.grad(loss_fn)(params, data) return loss_fn(params - lr * grads, data) meta_grads jax.grad(meta_loss_fn)(params, data)这里meta_loss_fn内部先调用jax.grad(loss_fn)得到一阶梯度再执行一步参数更新外层再用jax.grad对这个含梯度运算的函数求导。由于 JAX 的变换可以任意组合与嵌套meta_grads会正确地携带二阶信息对梯度更新路径的导数这正是 MAML 所需的效果——无需任何框架级的特殊支持。停止梯度jax.lax.stop_gradient自动微分会自动计算函数对输入的梯度但有时你需要更多控制例如避免梯度穿过计算图的一部分如只让某个损失影响网络的部分参数。源码层面的实现原理jax.lax.stop_gradient在 jax/_src/lax/lax.py 中实现。从实现看它操作上等价于恒等函数原样返回输入x但通过绑定ad_util.stop_gradient_p这一自动微分原语来阻断梯度流如果输入是扩展 dtype如字符串直接原样返回如果输入处于前向模式 JVP 追踪中则返回其 primal去掉切线否则绑定stop_gradient_p原语。该原语配套的 JVP 规则jax/_src/lax/lax.py将切线置零批量规则则透传从而在正向与反向模式中都不让梯度通过。官方文档特别指出如果有嵌套的多重梯度计算stop_gradient会同时阻断所有这些层级的梯度流。实战案例TD(0) 强化学习更新考虑 TD(0)时序差分强化学习更新用于从与环境交互的经验中学习状态的价值估计。假设价值估计 $v_{\theta}(s_{t-1})$ 由线性函数参数化# Value function and initial parameters value_fn lambda theta, state: jnp.dot(theta, state) theta jnp.array([0.1, -0.1, 0.])考虑从状态 $s_{t-1}$ 转移到 $s_t$ 并观察到奖励 $r_t$ 的一次转移# An example transition. s_tm1 jnp.array([1., 2., -1.]) r_t jnp.array(1.) s_t jnp.array([2., 1., 0.])TD(0) 对网络参数的更新为$$ \Delta \theta (r_t v_{\theta}(s_t) - v_{\theta}(s_{t-1})) \nabla v_{\theta}(s_{t-1}) $$注意这个更新并不是任何损失函数的梯度。但它可以写成如下伪损失函数的梯度$$ L(\theta) - \frac{1}{2} [r_t v_{\theta}(s_t) - v_{\theta}(s_{t-1})]^2 $$前提是忽略目标 $r_t v_{\theta}(s_t)$ 对参数 $\theta$ 的依赖。如果朴素地写出伪损失梯度会包含target对 $\theta$ 的依赖从而得到错误结果def td_loss(theta, s_tm1, r_t, s_t): v_tm1 value_fn(theta, s_tm1) target r_t value_fn(theta, s_t) return -0.5 * ((target - v_tm1) ** 2) td_update jax.grad(td_loss) delta_theta td_update(theta, s_tm1, r_t, s_t) delta_theta正确做法是用jax.lax.stop_gradient强制 JAX 忽略target对 $\theta$ 的依赖def td_loss(theta, s_tm1, r_t, s_t): v_tm1 value_fn(theta, s_tm1) target r_t value_fn(theta, s_t) return -0.5 * ((jax.lax.stop_gradient(target) - v_tm1) ** 2) td_update jax.grad(td_loss) delta_theta td_update(theta, s_tm1, r_t, s_t) delta_theta这样target被当作与 $\theta$无关的常数梯度计算得到正确的 TD(0) 参数更新。交叉验证用原始 TD(0) 更新表达式复核下面用原始的 TD(0) 更新表达式直接计算 $\Delta \theta$ 作为交叉验证建议读者先尝试用jax.grad自行实现s_grad jax.grad(value_fn)(theta, s_tm1) delta_theta_original_calculation (r_t value_fn(theta, s_t) - value_fn(theta, s_tm1)) * s_grad delta_theta_original_calculation # [1.2, 2.4, -1.2], same as delta_theta结果为[1.2, 2.4, -1.2]与delta_theta完全一致。jax.lax.stop_gradient在其他场景同样有用例如希望某个损失的梯度只影响神经网络的一部分参数其余参数由另一个损失训练时。直通估计器Straight-through estimator直通估计器是一种为本身不可微的函数定义梯度的技巧。给定不可微函数 $f : \mathbb{R}^n \to \mathbb{R}^n$ 作为更大函数的一部分我们希望在反向传播时假装 $f$ 是恒等函数。用jax.lax.stop_gradient可以优雅地实现def f(x): return jnp.round(x) # non-differentiable def straight_through_f(x): # Create an exactly-zero expression with Sterbenz lemma that has # an exactly-one gradient. zero x - jax.lax.stop_gradient(x) return zero jax.lax.stop_gradient(f(x)) print(f(x): , f(3.2)) print(straight_through_f(x):, straight_through_f(3.2)) print(grad(f)(x):, jax.grad(f)(3.2)) print(grad(straight_through_f)(x):, jax.grad(straight_through_f)(3.2))这段代码的精妙之处在于Sterbenz 引理的运用zero x - jax.lax.stop_gradient(x)在数值上精确等于零Sterbenz 引理保证相邻浮点数相减无舍入误差但其梯度恰为 1而jax.lax.stop_gradient(f(x))提供不可微函数的前向值 $f(x)$ 但梯度为 0。两者相加后前向值等于 $f(x)$梯度等于恒等函数的梯度。因此jax.grad(f)(3.2)为 0round几乎处处导数为 0而jax.grad(straight_through_f)(3.2)为 1.0。逐样本梯度Per-example gradients大多数机器学习系统为了计算效率与方差缩减从批量数据计算梯度。但某些场景需要批次中每个样本各自对应的梯度例如基于梯度大小对数据进行优先级排序或逐样本进行梯度裁剪/归一化。在许多框架如 PyTorch、TensorFlow、Theano中库会直接在批维度上累加梯度导致逐样本梯度并不容易计算而对每个样本分别算 loss 再聚合的朴素变通方案通常效率极低。在 JAX 中只需把jax.jit、jax.vmap、jax.grad三个变换组合起来即可高效实现perex_grads jax.jit(jax.vmap(jax.grad(td_loss), in_axes(None, 0, 0, 0))) # Test it: batched_s_tm1 jnp.stack([s_tm1, s_tm1]) batched_r_t jnp.stack([r_t, r_t]) batched_s_t jnp.stack([s_t, s_t]) perex_grads(theta, batched_s_tm1, batched_r_t, batched_s_t)逐层拆解这个组合第一步jax.grad得到单样本梯度函数。对td_loss应用jax.grad得到对单未批量化输入计算参数梯度的函数dtdloss_dtheta jax.grad(td_loss) dtdloss_dtheta(theta, s_tm1, r_t, s_t)该函数计算出上面数组中的一行。第二步jax.vmap向量化。对单样本梯度函数应用jax.vmap会给所有输入输出增加一个批维度。给定一批输入即产生一批输出——每个输出对应输入批次中对应样本的梯度almost_perex_grads jax.vmap(dtdloss_dtheta) batched_theta jnp.stack([theta, theta]) almost_perex_grads(batched_theta, batched_s_tm1, batched_r_t, batched_s_t)但这样还不太对我们被迫手动传入一批theta而我们实际只想用单个theta。通过给jax.vmap添加in_axes参数修复将theta指定为None其余参数指定为0使结果函数只对其余参数增加额外轴而theta保持未批量化inefficient_perex_grads jax.vmap(dtdloss_dtheta, in_axes(None, 0, 0, 0)) inefficient_perex_grads(theta, batched_s_tm1, batched_r_t, batched_s_t)这一步功能正确但还没有发挥 JAX 的编译优势。第三步jax.jit编译加速。将整个函数用jax.jit包裹得到编译后的高效版本perex_grads jax.jit(inefficient_perex_grads) perex_grads(theta, batched_s_tm1, batched_r_t, batched_s_t)用%timeit对比两个版本的耗时配合.block_until_ready()确保等待异步执行完成得到真实的计时结果%timeit inefficient_perex_grads(theta, batched_s_tm1, batched_r_t, batched_s_t).block_until_ready() %timeit perex_grads(theta, batched_s_tm1, batched_r_t, batched_s_t).block_until_ready()jit版本会在首次调用时完成编译之后每次调用都直接执行编译后的计算显著快于未编译版本。Hessian-vector productgrad-of-grad利用高阶jax.grad可以构造 Hessian-vector productHVP函数。HVP 在截断牛顿共轭梯度算法用于最小化光滑凸函数中很有用也可用于研究神经网络训练目标的曲率。对具有连续二阶导数Hessian 对称的标量值函数 $f : \mathbb{R}^n \to \mathbb{R}$Hessian 记为 $\partial^2 f(x)$HVP 函数计算$$v \mapsto \partial^2 f(x) \cdot v$$对任意 $v \in \mathbb{R}^n$。关键技巧不要实例化完整的 Hessian 矩阵。在神经网络场景中 $n$ 可能高达百万甚至数十亿存储完整 Hessian 矩阵完全不可行。幸运的是利用如下恒等式即可写出高效的 HVP$$\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$ 做点积。注意我们始终只对向量值自变量的标量值函数求导这正是jax.grad最高效的形态。JAX 代码只需寥寥数行def hvp(f, x, v): return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)这个例子还展示了 JAX 的一个特性可以自由使用词法闭包lexical closureJAX 的追踪器不会被闭包捕获的外部变量所困扰。文末我们会用稠密 Hessian 验证该实现并给出混合前向/反向模式forward-over-reverse的更高效版本。实现说明内层jax.grad(f)(x)计算梯度 $\partial f(x)$jnp.vdot与 $v$ 做点积得到标量 $g(x)$外层jax.grad对 $g$ 求梯度即得 $\partial^2 f(x) v$。整个过程只需要 $O(n)$ 的额外内存而非 Hessian 的 $O(n^2)$。用 jacfwd 与 jacrev 计算完整 Jacobian 与 Hessian可以用jax.jacfwd和jax.jacrev计算完整的 Jacobian 矩阵from jax import jacfwd, jacrev # Define a sigmoid function. def sigmoid(x): return 0.5 * (jnp.tanh(x / 2) 1) # Outputs probability of a label being true. def predict(W, b, inputs): return sigmoid(jnp.dot(inputs, W) b) # Build a toy dataset. 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]]) # Initialize random model coefficients 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, ()) # Isolate the function from the weight matrix to the predictions 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通常略占优势。支持容器类型pytreesjax.jacfwd与jax.jacrev还支持容器类型参数def predict_dict(params, inputs): return predict(params[W], params[b], inputs) J_dict jax.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)返回的 Jacobian 结构与输入 pytree 结构对应每个叶子给出对应参数到输出的 Jacobian 块。组合两者计算稠密 Hessian组合使用jax.jacfwd与jax.jacrev即可得到稠密 Hessiandef hessian(f): return jax.jacfwd(jax.jacrev(f)) H hessian(f)(W) print(hessian, with shape, H.shape) print(H)形状规律为什么输出是 m×n×n这个形状符合预期。若从函数 $f : \mathbb{R}^n \to \mathbb{R}^m$ 出发在点 $x \in \mathbb{R}^n$ 处应得到如下形状$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。依此类推更高阶导数会继续叠加输入维度。为什么 forward-over-reverse 通常最优实现hessian理论上可以用jacfwd(jacrev(f))、jacrev(jacfwd(f))或其他任意组合。但前向套反向forward-over-reverse通常最高效。原因在于内层 Jacobian 计算时我们通常在对一个宽扁 Jacobian 的函数求导比如损失函数 $f : \mathbb{R}^n \to \mathbb{R}$输入远多于输出正是反向模式的主场而外层 Jacobian 计算时我们在对具有方形 Jacobian 的函数求导因为 $\nabla f : \mathbb{R}^n \to \mathbb{R}^n$输入输出维度相同这恰是前向模式占优的场景。源码印证jax.jacfwd与jax.jacrev均定义于 jax/_src/api.py。其中jacfwd的实现路径是先用argnums_partial2固定非求导参数对每个标准基向量调用_jvp即前向模式的 JVP再用jax.vmap批量执行以逐列填充 Jacobianjax/_src/api.py同时还会做输入输出 dtype 检查——非全纯模式下要求实值浮点输入。jax.hessian在 jax/_src/api.py 中实现其定义正是jacfwd(jacrev(fun, ...), ...)即前向套反向组合并原生支持 pytree 输入输出Hessian 的树结构由输出树结构与输入树结构的两份拷贝的树积构成每个叶子块形状为(out..., in1..., in2...)。若需将 pytree 展平为 1D 向量可参考jax.flatten_util.flatten_pytree。总结本文围绕 JAX 高阶导数这一主题覆盖了从基础到进阶的完整知识链技术点核心实现典型场景Hessian稠密jax.jacfwd(jax.grad(f))或内置jax.hessian即jacfwd(jacrev(f))二阶优化、曲率分析元学习对含梯度更新的函数再次jax.gradMAML 等梯度截断jax.lax.stop_gradient底层绑定stop_gradient_p原语TD(0) 强化学习、分参数训练直通估计器zero jax.lax.stop_gradient(f(x))Sterbenz 引理构造量化/离散化模块的反向传播逐样本梯度jax.jit(jax.vmap(jax.grad(f), in_axes(None, 0, ...)))按梯度优先级采样、逐样本裁剪Hessian-vector productjax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)截断牛顿共轭梯度、曲率研究Jacobianjax.jacfwd前向适合 tall/jax.jacrev反向适合 wide灵敏度分析、pytree 参数微分贯穿始终的核心思想是JAX 中任何求导变换本身都是可微分的函数因此高阶导数不过是变换的叠加。无论是二阶的 Hessian、需要截断梯度流的 TD(0) 更新、还是融合编译与向量化的逐样本梯度都能用少量标准变换的嵌套组合直接表达这正是 JAX 在元学习与科研计算中被广泛使用的原因。若希望深入前向/反向模式的底层原理以及jacfwd/jacrev的高效实现细节可继续阅读 自动微分教程 与 advanced_autodiff.md。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表