- 机器学习
- 深度学习
【免费下载链接】jax
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
导读
本文基于 JAX 仓库中的 JEP 文档 docs/jep/4008-custom-vjp-update.md,系统讲解jax.custom_vjp与nondiff_argnums的使用边界在 PR #4008 之后发生的关键变化:Tracer 不再允许传入nondiff_argnums位置,数组型非可微参数应改为普通参数配合None占位,同时custom_jvp/custom_vjp对 Tracer 的词法闭包问题被彻底修复。读者学完后,能够正确迁移旧式custom_vjp代码、理解何时仍必须使用nondiff_argnums,并掌握底层实现与测试证据,避免在新代码中踩坑。
背景:custom_vjp与nondiff_argnums是什么
jax.custom_vjp是 JAX 中为函数自定义反向模式微分(VJP)规则的装饰器。它要求用户提供一对规则:
- fwd 规则:输入与原始函数相同,输出一个二元组
(primal_out, residuals),其中residuals是在前向传播中保存、供反向传播使用的值; - bwd 规则:接收
residuals与输出余切(cotangent)g,返回一个长度等于原始函数参数个数的元组,代表各参数的梯度。
装饰器还接受可选的nondiff_argnums参数,用于标记"不可微参数"的位置。历史上,它被用来声明某些参数不需要梯度,例如标量阈值、控制标志等。完整接口定义位于 jax/_src/custom_derivatives.py 中的custom_vjp类(约第 459 行起),其defvjp方法负责注册 fwd/bwd 规则并支持symbolic_zeros选项。
关于custom_vjp的入门教程可参考仓库内 docs/notebooks/Custom_derivative_rules_for_Python_code.ipynb(及同名 .md 版本),本文默认读者已熟悉其基本用法。
核心变更:nondiff_argnums不再接受 Tracer
变更内容
JAX PR #4008 之后,传入custom_vjp函数nondiff_argnums位置的参数不能是 Tracer(或包含 Tracer 的容器)。所谓 Tracer,是 JAX 在jit、vmap、grad等变换中用来追踪计算的抽象对象——只要参数在某个变换内部流动,它就是 Tracer。
这一限制的本质含义是:
- 数组型参数不应放入
nondiff_argnums; nondiff_argnums只应保留给非数组值,例如 Python 可调用对象(函数)、shape 元组、字符串等。
迁移规则非常简单:凡是在旧代码里把数组值放进nondiff_argnums的地方,直接把它当作普通参数传递;在bwd规则中,对这些参数的位置返回None,表示"没有对应的梯度值"。
旧写法为什么不再可靠
以下是 JEP 文档给出的旧式clip_gradient写法,它把lo、hi两个数组放在nondiff_argnums=(0, 1):
from functools import partial import jax import jax.numpy as jnp @partial(jax.custom_vjp, nondiff_argnums=(0, 1)) def clip_gradient(lo, hi, x): return x # identity function def clip_gradient_fwd(lo, hi, x): return x, None # no residual values to save def clip_gradient_bwd(lo, hi, _, g): return (jnp.clip(g, lo, hi),) clip_gradient.defvjp(clip_gradient_fwd, clip_gradient_bwd)这段代码在lo或hi来自jit/vmap/grad等变换(即成为 Tracer)时无法工作——这正是 PR #4008 要消除的缺陷。
新写法:数组参数走常规通道
迁移后的新写法不再使用nondiff_argnums,lo、hi作为普通参数参与,在前向规则中作为 residual 保存,在反向规则中返回None占位:
import jax import jax.numpy as jnp @jax.custom_vjp # no nondiff_argnums! def clip_gradient(lo, hi, x): return x # identity function def clip_gradient_fwd(lo, hi, x): return x, (lo, hi) # save lo and hi values as residuals def clip_gradient_bwd(res, g): lo, hi = res return (None, None, jnp.clip(g, lo, hi)) # return None for lo and hi clip_gradient.defvjp(clip_gradient_fwd, clip_gradient_bwd)注意两个关键变化:
- fwd 签名与原始函数完全一致,且必须返回二元组
(primal, residuals)——源码 custom_derivatives.py 中defvjp的文档(约第 513–521 行)明确要求 fwd 返回"primal output + residual"的 pair,若结构不符会抛出带详细说明的TypeError; - bwd 只接收两个参数
(res, g),res是 fwd 保存的 residual 元组,返回值中lo、hi对应的位置填None,最后一个位置才是x的梯度。
如果仍沿用旧写法,JAX 会在任何可能出错的场景(即有 Tracer 传入nondiff_argnums)抛出清晰、响亮的错误,而不是静默产生错误结果。
何时仍必须使用nondiff_argnums
并非所有nondiff_argnums用法都应废弃。当参数本身是无法作为 JAX 值参与变换的对象(典型如 Python 函数)时,nondiff_argnums依然是唯一正确的手段。JEP 文档给出了skip_app示例:第一个参数是被调用函数f,它是不可微、也无法成为 Tracer 的非数组值:
from functools import partial import jax @partial(jax.custom_vjp, nondiff_argnums=(0,)) def skip_app(f, x): return f(x) def skip_app_fwd(f, x): return skip_app(f, x), None def skip_app_bwd(f, _, g): return (g,) skip_app.defvjp(skip_app_fwd, skip_app_bwd)注意此处bwd的签名是(f, _, g):f作为nondiff_argnums参数被以特殊的前置参数形式传入 bwd,这与新式写法中 residual 从元组里解包不同。在 custom_derivatives.py 的实现里,nondiff_argnums路径通过argnums_partial拆分出static_args,再用_add_args把静态参数拼接回 bwd 规则的前端(约第 603–612 行),这正是"特殊前置参数"语义的来源。
底层原理:为什么曾经是 bug,现在如何被守卫
本质:nondiff_argnums曾像词法闭包一样工作
JEP 文档指出,nondiff_argnums旧实现的运作方式"非常像词法闭包(lexical closure)"——即把被标记的参数静默地绑定到规则内部,如同被闭包捕获的变量。然而在 PR #4008 之前,custom_jvp/custom_vjp并不支持对 Tracer 做词法闭包。部分场景碰巧能工作,但更多场景会抛出复杂且令人困惑的错误信息。设计上的这一失误正是缺陷根源。
PR #4008:修复词法闭包
PR #4008 修复了custom_jvp与custom_vjp的全部词法闭包问题。修复之后:
- 对所有非自动微分变换(如
jit、vmap),在custom_jvp/custom_vjp函数或规则中闭包捕获 Tracer 都会"直接可用"(Just Work); - 对自动微分变换(如
grad),若试图对被闭包捕获的值求导,会得到一段明确说明原因的错误信息:
Detected differentiation of a custom_jvp function with respect to a closed-over value. That isn't supported because the custom JVP rule only specifies how to differentiate the custom_jvp function with respect to explicit input parameters.
Try passing the closed-over value into the custom_jvp function as an argument, and adapting the custom_jvp rule.
(大意:检测到对 custom_jvp 函数被闭包捕获的值求导,这不受支持,因为自定义 JVP 规则只说明了如何对显式输入参数求导;请把闭包值作为参数传入并调整规则。)
为什么选择"禁止"而不是"支持"
JEP 文档解释了设计取舍:若允许custom_vjp的nondiff_argnums接收 Tracer,需要做大量簿记工作——重写用户的 fwd 规则把值作为 residual 返回,并重写 bwd 规则把它们当作普通 residual 接收(而非nondiff_argnums式的特殊前置参数)。这还要处理任意 pytree 结构,复杂度高且没必要:只要用户把数组型不可微参数当作普通参数和 residual 处理,一切已经正常运作。
源码中的守卫:_check_for_tracers
当前仓库 jax/_src/custom_derivatives.py 中,custom_vjp.__call__会在使用nondiff_argnums时对每个指定位置的参数调用_check_for_tracers(约第 603–604 行)。该函数遍历 pytree 叶子,一旦发现core.Tracer就抛出UnexpectedTracerError(第 652–662 行):
def _check_for_tracers(x): for leaf in tree_leaves(x): if isinstance(leaf, core.Tracer): msg = ("Found a JAX Tracer object passed as an argument to a custom_vjp " "function in a position indicated by nondiff_argnums as " "non-differentiable. Tracers cannot be passed as non-differentiable " "arguments to custom_vjp functions; instead, nondiff_argnums should " "only be used for arguments that can't be or contain JAX tracers, " "e.g. function-valued arguments. In particular, array-valued " "arguments should typically not be indicated as nondiff_argnums.") raise UnexpectedTracerError(msg)这条报错信息非常具有指导性:只有在参数"不可能成为或包含 JAX tracer"时才适合用nondiff_argnums,数组型参数则"通常不应"标记为nondiff_argnums。
与custom_jvp的差异:为何只有custom_vjp需要迁移
JEP 文档特别指出:与custom_vjp不同,让custom_jvp的nondiff_argnums参数接收 Tracer实现起来很容易,因此本次迁移只针对custom_vjp。
从源码可以印证这一不对称性。在 custom_derivatives.py 中,custom_jvp的__call__(约第 245–253 行)对nondiff_argnums位置的参数直接套用了_stop_gradient:
if self.nondiff_argnums: nondiff_argnums = set(self.nondiff_argnums) args = tuple(_stop_gradient(x) if i in nondiff_argnums else x for i, x in enumerate(args)) diff_argnums = [i for i in range(len(args)) if i not in nondiff_argnums] f_, dyn_args = argnums_partial(lu.wrap_init(self.fun), diff_argnums, args, require_static_args_hashable=False) static_args = [args[i] for i in self.nondiff_argnums] jvp = _add_args(lu.wrap_init(self.jvp), static_args)对 Tracer 施加stop_gradient是 JAX 中的成熟操作,天然可行;而custom_vjp走的是custom_vjp_call_p原语绑定路径,需要对 fwd/bwd 规则做结构改写,成本完全不同。这就是"更新只发生在custom_vjp侧"的实现原因。
整数参数的额外利好(PR #4039)
JEP 文档还预告了 PR #4039 带来的改进:在 #4039 之前,JAX 在自动微分中遇到整型输入/输出时可能报错;而 #4039 之后,整型输入输出参与 autodiff 也能直接工作。这进一步降低了"把数组型不可微参数当作普通参数 + None 占位"迁移方案的摩擦——例如clip_gradient中若lo/hi是整型边界,新写法同样成立。仓库测试 tests/api_test.py 中的test_nondiff_arg(约第 8329 行)正是"函数作为nondiff_argnums参数"的合法用法验证:它用lambda x: 2 * x作为第一个参数,并对x正常求value_and_grad,梯度结果为jnp.cos(1.),说明函数型不可微参数与数组参数混用完全正常。
测试证据:错误被显式守卫,闭包被显式支持
仓库 tests/api_test.py 中围绕本次变更有一组针对性测试,可作为迁移行为的事实依据:
test_nondiff_arg_tracer_error(约第 8419 行):定义@partial(jax.custom_vjp, nondiff_argnums=(0,))的函数后,在@jit包装下调用(参数为 Tracer),断言抛出UnexpectedTracerError且消息包含"custom_vjp"。这验证了"旧写法会得到响亮错误"的设计目标;test_closed_over_jit_tracer(约第 8347 行):原本测试"jit 闭包捕获 Tracer"的场景,现已被SkipTest跳过,注释明确指出"该行为不再被支持",理由是"禁止nondiff_argnums中的 Tracer 以大幅简化簿记,同时仍支持必要的场景";- **
test_closed_over_vmap_tracer(约第 8378 行)**与test_closed_over_tracer3(约第 8398 行):验证修复后custom_vjp函数/规则可以闭包捕获 Tracer 并在vmap下正常工作,甚至可以对闭包捕获的值参与反向传播(test_closed_over_tracer3将x放入 residual 并从 bwd 中使用); - **
test_closure_convert(约第 8982 行)**与test_closure_convert_mixed_consts(约第 9019 行):展示jax.closure_convert与函数型nondiff_argnums参数配合的实际模式——先closure_convert把闭包转成显式 aux 参数,再交给nondiff_argnums=(0,)的custom_vjp函数处理,支持对c、x等多个参数同时求梯度(对应梯度分别为42. * c与17. * x)。
迁移自检清单
完成代码迁移后,可用以下清单快速自检:
- 审查每个
nondiff_argnums位置:参数是否可能是数组或含数组的 pytree?若是,改为普通参数; - fwd 规则签名:应与原始函数参数完全一致,输出
(primal, residuals)二元组,需要保留下来的非可微数组必须存入 residual; - bwd 规则签名:改为
(res, g)形式(无nondiff_argnums时),按位置返回元组,长度等于参数个数,不可微数组位置填None; - 保留
nondiff_argnums的场景:仅用于函数、shape 元组、字符串等不可能成为 Tracer 的非数组值,此时 bwd 仍以"特殊前置参数"接收它们; - 验证:在
jit、vmap、grad组合下运行并核对梯度数值,可参考 tests/api_test.py 中对应测试的断言方式。
总结
PR #4008 更新确立了custom_vjp的一项明确约定:nondiff_argnums只服务于非数组静态值,数组型不可微参数一律走"普通参数 + residual 保存 +None梯度占位"的常规通道。这一约定消除了nondiff_argnums旧实现中"类闭包绑定 Tracer"的缺陷,换来了对所有变换的稳健支持;对于误用,源码中的_check_for_tracers会抛出带迁移建议的明确错误;而词法闭包捕获 Tracer 的能力(尤其配合vmap)则在测试中被充分验证。JAX 开发者可依据本文对照 docs/jep/4008-custom-vjp-update.md 与 jax/_src/custom_derivatives.py 源码,完成存量代码的安全迁移。
- 机器学习
- 深度学习
【免费下载链接】jax
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
相关推荐
JAX `custom_vjp` 与 `nondiff_argnums` 升级指南:从闭包式非可微参数迁移到残差式自定义 VJP
JAX custom_vjp 与 nondiff_argnums 升级指南:从闭包式非可微参数迁移到残差式自定义 VJP 本篇指南聚焦 JAX 增强提案 JEP
人工智能机器学习深度学习编译器高性能计算ggplot2版本更新详解:3.5.0新特性与迁移指南
ggplot2版本更新详解:3.5.0新特性与迁移指南 ggplot2作为R语言中最受欢迎的数据可视化包,在3.5.0版本中带来了令人兴奋的新功能和改进。🎉
数据可视化Prophet 迁移指南:fbprophet 兼容层(shim)与 v1.0 包名更迭详解
Prophet 迁移指南:fbprophet 兼容层(shim)与 v1.0 包名更迭详解 导读 本指南面向所有在 Python 时间序列预测项目中使用过 fb
数据分析
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考