深入理解 JAX 自定义导数规则:custom_jvp / custom_vjp 设计原理与实现剖析
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
本文以 JAX 仓库中的设计文档 docs/jep/2026-custom-derivatives.md 为主体,系统梳理jax.custom_jvp与jax.custom_vjp的设计动机、核心问题、解决方案与实现机制。你将理解为什么旧版custom_transforms机制会在vmap组合下丢失自定义导数规则、如何通过core.call风格的 Python 级调用原语解决该语义问题,以及 JVP/VJP 规则的具体 API 约定(含nondiff_argnums、kwargs、pytree 支持等)。文中结合仓库源码(jax/_src/custom_derivatives.py、tests/custom_api_test.py、jax/experimental/ode.py)给出实现级证据与可复现示例,帮助读者在自定义数值稳定导数、调试 NaN、实现odeint式高阶函数时正确使用这套机制。
背景:JAX 中定义求导规则的两种途径
在 JAX 中,有两种方式可以定义求导规则:
- 使用
jax.custom_jvp与jax.custom_vjp,为已经是 JAX 可变换(JAX-transformable)的 Python 函数定义自定义的前向(JVP)与反向(VJP)求导规则。这是本文的绝对主题。 - 定义新的
core.Primitive实例并为其实现全部变换规则,例如接入求解器、仿真器等外部数值系统的函数调用。这是更高阶、更低层的方式,本文不展开。
作为 JAX 开发者,我们希望以logit、expit(见 jax/scipy/special.py 中的相关实现)为代表,用其他原语定义库函数,但在求导时表现出"原语级"行为——即显式给出可能更数值稳定或性能更好的自定义导数规则,同时不必为这些函数单独指定vmap、jit规则。作为长远目标(stretch goal),还希望让高阶函数(fixed_point、odeint、root等)能方便地接入自定义求导规则。
设计文档明确指出本设计要解决的问题清单:关闭 #116、#1097、#1249、#1275 等一系列 issue,并取代旧的custom_transforms机制。
设计目标(Goals)与非目标(Non-goals)
核心目标
设计文档将目标明确划分为两级:
- 用户侧:希望用户能自定义其代码的前向/反向求导行为,且该自定义
- 具备清晰一致的语义,能正确与其他 JAX 变换组合;
- 足够灵活,支持 Autograd、PyTorch 中常见的用法,包括对 Python 控制流求导、NaN 调试等场景。
- 开发者侧:
logit/expit这类以原语组合定义的库函数,求导时具有"原语级"自定义规则,而无需额外提供vmap/jit规则。
由此归纳出的主要目标是:
- 解决vmap 移除自定义 JVP 语义问题(对应 issue #1249);
- 允许在自定义 VJP 中使用 Python(例如调试 NaN,对应 issue #1275)。
次要目标包括:简化用户体验(符号零、kwargs 等);推动用户能轻松为fixed_point、odeint、root等添加自定义规则。
明确排除的非目标
- 不做变换泛化的自定义机制:旧的
custom_transforms目标是"变换泛化"地自定义行为,理论上允许用户为任意变换定制规则。本设计只解决求导(JVP 与 VJP 分开)的自定义问题——这是实际唯一被请求的场景;通过专精化降低了复杂度、提升了灵活性。若要控制所有规则,直接编写原语即可。 - 不优先追求数学美学:虽然自定义 VJP 的签名
a -> (b, CT b --o CT a)在数学上更优雅,但 Python 机制难以处理返回类型中的闭包,因此实现上**显式处理残差(residuals)**而非依赖闭包。 - 暂不支持序列化:将 staged-out 的序列化程序表示加载后继续做 JAX 变换(而不只是求值)目前不在范围内。这为将来把 Python 可调用对象"藏"在哪里保留了灵活性。
两个核心问题
问题一:vmap 移除自定义 JVP 语义问题
这是本设计文档最核心的动机。旧custom_transformsAPI 存在一个反直觉的 bug:
# 旧 custom_transforms API(将被替换) @jax.custom_transforms def f(x): return 2. * x # f_vjp :: a -> (b, CT b --o CT a) def f_vjp(x): return f(x), lambda g: 3. * x # 3 而不是 2 jax.defvjp_all(f, f_vjp) grad(f)(1.) # 3. vmap(grad(f))(np.ones(4)) # [3., 3., 3., 3.] grad(lambda x: vmap(f)(x).sum())(np.ones(4)) # [2., 2., 2., 2.] ← 意外!最后一行grad套vmap的结果不符合预期。一般来说,施加vmap(或任何非求导变换)都会"移除"自定义求导规则(施加jvp时,若定义了自定义 VJP 规则则会直接报错)。
根源分析
变换本质上是"重写"(rewrites)。custom_transforms机制会让求值f(x)时应用如下 jaxpr:
{ lambda ; ; a. let b = f_primitive a in [b] }其中f_primitive是为每个custom_transforms函数(实际上每次调用都会)新引入的原语,自定义 VJP 规则就挂在该原语上。求grad(f)(x)时,微分机制遇到f_primitive便用自定义规则处理。
然而,f_primitive对vmap是透明的:vmap相当于内联(inlining)f_primitive的定义,于是vmap(f)实际变成:
{ lambda ; ; a. let b = mul 2. a in [b] }即vmap把函数重写为其底层原语组合,完全移除了f_primitive,自定义规则随之丢失。
语义不一致性
更一般地,因为vmap(f)(xs) == np.stack([f(x) for x in xs])是vmap的语义定义,那么必须成立:
jvp(vmap(f))(xs) == jvp(lambda xs: np.stack([f(x) for x in xs]))但当f定义了自定义导数规则时,该性质不成立:右侧使用了自定义规则,左侧却没有。
设计文档强调:该问题不限于vmap,凡是"变换一个函数f的语义由对f的调用来定义(而非重写为另一个函数)"的变换都会受影响,mask变换也属于此类;而求导变换不在其列。如果再加上自定义vmap规则等,交互会愈发复杂——这正说明custom_transforms的"变换泛化"问题框架过于宽泛。
问题二:Python 灵活性缺失
与 Autograd、PyTorch 相同(但不同于 TF1),JAX 对 Python 函数的求导是在函数执行与追踪(tracing)的同时进行的。这带来两个关键好处:
- 支持基于 pdb 的工作流:用户可以用标准 Python 调试器检查数值、捕获 NaN。设计文档作者特别提到,在实现
odeint原语期间多次依靠运行时数值检查来调试问题;一个特别实用的技巧是在自定义 VJP 规则中插入调试器断点,从而在反向传播的特定位置进入调试器。 - 允许对 Python 原生控制流求导:
if x > 0这样的原生分支可以直接参与求导。
但旧custom_transforms机制做不到这点:因为它对用户函数和自定义规则都预先形成 jaxpr,遇到 Python 控制流就会报抽象值追踪错误:
# 旧 custom_transforms API(将被替换) @jax.custom_transforms def f(x): if x > 0: return x else: return 0. def f_vjp(x): return ... jax.defvjp_all(f, f_vjp) grad(f)(1.) # Error!解决方案:借鉴 core.call 的 custom_jvp_call 原语
设计文档的核心思想非常简洁:core.call已经解决了这些问题。将"为用户函数指定自定义 JVP 规则"这一任务,表述为一种新的 Python 级调用原语——记为custom_jvp_call(注意:不加入 jaxpr 语言本身)。
custom_jvp_call与core.call一样关联一个用户 Python 函数,但额外携带第二个 Python 可调用对象表示 JVP 规则:
vmap(call(f)) == call(vmap(f)) # core.call 的行为 vmap(custom_jvp_call(f, f_jvp)) == custom_jvp_call(vmap(f), vmap(f_jvp))即vmap等变换对custom_jvp_call直接"穿过",施加到底层的两个 Python 可调用对象上。这从机制上解决了 vmap 移除自定义 JVP 语义问题。
jvp变换的交互则符合直觉——直接调用f_jvp:
jvp(call(f)) == call(jvp(f)) jvp(custom_jvp_call(f, f_jvp)) == f_jvp而求值与编译是"退出" JAX 系统的两种方式(之后不再有变换施加),规则平凡:
eval(call(f)) == eval(f) jit(call(f)) == hlo_call(jit(f)) eval(custom_jvp_call(f, f_jvp)) == eval(f) jit(custom_jvp_call(f, f_jvp)) == hlo_call(jit(f))即若 JVP 规则尚未把custom_jvp_call(f, f_jvp)重写为f_jvp,到达求值或jit阶段时求导不会再发生,直接忽略f_jvp、表现得像core.call即可。
唯一的小坑:initial-style 原语与动态作用域
lax.scan这类"initial-style"jaxpr 形成原语是个例外:它对 jaxpr 的"staging out"不退出变换系统——对lax.scan施加 jvp 或 vmap 时,需要把它施加到 jaxpr 所表示的函数上。因此这类原语依赖"jaxpr ↔ Python 可调用对象"的往返且保持语义不变,自定义导数规则的语义也必须被保留。
解决方案是利用一点动态作用域:当为 initial-style 原语(如 jax/_src/lax/control_flow.py 中的机制)staging out jaxpr 时,在全局追踪状态上设置一个标志位。该标志置位时,不再使用 final-style 的custom_jvp_call,而是改用 initial-style 的custom_jvp_call_jaxpr原语,并预先将f与f_jvp追踪为 jaxpr,以简化 initial-style 处理。
脚注:虽然道德上应在绑定
custom_jvp_call_jaxpr前就为f与f_jvp都形成 jaxpr,但必须延迟f_jvp的 jaxpr 形成——因为它可能调用自定义 JVP 函数本身,提前处理会导致无限递归;延迟方式是把 jaxpr 形成放进一个 thunk(惰性闭包)中。
设计文档还指出:如果放弃 Python 灵活性目标,只保留custom_jvp_call_jaxpr、去掉单独的 Python 级custom_jvp_call也够用——但正是这份"Python 灵活性"让 final-style 原语不可或缺。
API 约定:JVP 与 VJP 规则的标准写法
自定义 JVP:(a, Ta) -> (b, T b)
对a -> b的函数,自定义 JVP 用一个(a, Ta) -> (b, T b)的函数指定(T表示切向量/tangent):
# f :: a -> b @jax.custom_jvp def f(x): return np.sin(x) # f_jvp :: (a, T a) -> (b, T b) def f_jvp(primals, tangents): x, = primals t, = tangents return f(x), np.cos(x) * t f.defjvp(f_jvp)关于高阶求导的关键提示:为了让规则能应用于高阶求导,必须在f_jvp的函数体内调用f(即上面示例中的return f(x), ...而非直接重算np.sin(x))。这一要求会排除f内部与切向量计算之间的某些"工作量共享"。
自定义 VJP:前向a -> (b, c)+ 反向(c, CT b) -> CT a
对a -> b的函数,自定义 VJP 拆成前向与反向两个函数:前向a -> (b, c)输出主值b与残差c;反向(c, CT b) -> CT a消费残差与输出余切(cotangent),产生输入余切:
# f :: a -> b @jax.custom_vjp def f(x): return np.sin(x) # f_fwd :: a -> (b, c) def f_fwd(x): return f(x), np.cos(x) # f_bwd :: (c, CT b) -> CT a def f_bwd(cos_x, g): return (cos_x * g,) f.defvjp(f_fwd, f_bwd)设计文档解释了为何不采用数学上更优雅的a -> (b, CT b --o CT a)签名:Python 可调用对象本质不透明(除非预先急切地追踪成 jaxpr,而那会带来表达力约束),且前向阶段可能返回一个闭包内含有vmaptracer 的可调用对象,实现复杂且可能牺牲表达性。因此选择显式残差方案。
其余 API 细节(bells and whistles)
- 任意 pytree:输入与输出类型
a、b、c可以是任意 jaxtypes 的 pytree(嵌套的 tuple/list/dict 等)。 - 按名传参(kwargs):当 kwargs 能通过
inspect模块解析为位置参数时支持按名传参。设计文档称其为一次"对 Python 3 增强的程序化签名检查能力的实验",可靠但不完备。 nondiff_argnums标记不可微参数:与jit的static_argnums类似,这些参数不要求是 JAX 类型。其传递约定为:设主函数签名为(d, a) -> b(d为不可微类型),则- JVP 规则签名为
(a, T a, d) -> T b——不可微参数按顺序排在primals与tangents之后; - VJP 规则的反向函数签名为
(d, c, CT b) -> CT a——不可微参数按顺序排在残差之前。
- JVP 规则签名为
当前仓库源码中的 API 佐证
仓库实现 jax/_src/custom_derivatives.py 与设计文档完全对应:
CustomJVPCallPrimitive(custom_jvp_call_p)与CustomVJPCallPrimitive(custom_vjp_call_p)分别对应设计中的两个 Python 级调用原语(见 jax/_src/custom_derivatives.py 与 jax/_src/custom_derivatives.py)。- 类
CustomJVP的__init__中,nondiff_argnames会通过fun_signature与infer_argnums_and_argnames解析为nondiff_argnums(jax/_src/custom_derivatives.py);defjvp支持symbolic_zeros选项,用于向规则传入静态符号零对象(jax/_src/custom_derivatives.py)。 defjvps是文档中提到的"便捷包装"的落地方案:为每个参数分别定义偏导规则(jax/_src/custom_derivatives.py)。注意文档倾向"最小化",因此便捷层收敛为defjvps这一种形式,且不可与nondiff_argnums混用。- 调用时,
__call__用resolve_kwargs(self.fun, args, kwargs)解析 kwargs,随后将函数与规则展平(_flatten_fun_nokwargs/_flatten_jvp),最终custom_jvp_call_p.bind(...)(jax/_src/custom_derivatives.py)。nondiff_argnums的实现在绑定前先对这些参数施加_stop_gradient并剥离(argnums_partial、prepend_static_args),从源码结构看与设计文档的"不可微参数单独传递"约定一致。 - 设计文档提到的
custom_jvp_call_jaxpr在源码中仍以兼容 stub 形式存在(jax/_src/custom_derivatives.py),注释说明其"仅为避免破坏内部用户而保留"。
测试用例印证
tests/custom_api_test.py 提供了与设计文档一一对应的验证:
- vmap 组合正确性(
test_vmap,tests/custom_api_test.py):分别验证vmap(f)、vmap(jvp(f))、jvp(vmap(f))、vmap(jvp(vmap(f)))的结果都与期望一致——这正是设计文档"vmap 穿过 custom_jvp_call"语义的直接测试。 nondiff_argnums/nondiff_argnames(tests/custom_api_test.py):示例将函数作为不可微参数传入jax.custom_jvp(nondiff_argnums=(0,)),JVP 规则签名为(f, primals, tangents),与文档约定的"不可微参数排在 primals/tangents 之后"完全吻合。
实现笔记:VJP 的两段式处理与 custom_lin
仓库实现印证了设计文档中的几个关键实现决策:
每个变换都有自定义绑定方法:为
custom_jvp_call与custom_vjp_call提供了类似core.call_bind的自定义 bind 方法,区别在于不处理 env traces(那些直接报错)。custom_lin原语与两段式反向求导:JAX 的反向自动微分被分解为**线性化(linearization)→ 部分求值(partial evaluation)→ 转置(transposition)**三步,因此自定义 VJP 规则分两步被处理:- 线性化步骤:
custom_vjp_call的 JVP 规则对切向量施加custom_lin; - 转置步骤:
custom_lin原语携带用户的 backward-pass 函数,作为原语只实现 transpose 规则。
在 jax/_src/custom_derivatives.py 中可以观察到配套约束:
custom_lin对vmap与 MLIR lowering 都直接抛错(raise_custom_vjp_error_on_jvp),即对自定义 VJP 函数施加jvp是不被允许的——这与设计文档"施加 jvp 时若定义了自定义 VJP 规则会失败"的表述一致。此外,对custom_vjp函数施加前向自动微分会触发disallow_jvp(jax/_src/custom_derivatives.py)。- 线性化步骤:
odeint作为规范用户案例:设计文档专门修订了jax.experimental.odeint,将其作为新 API 的"金标准用户"来检验 API 质量,并顺带做了三项改进:去掉 ravel/unravel 样板代码、用lax.scan替代索引更新逻辑、在简单摆锤基准上提速 20% 以上。当前仓库 jax/experimental/ode.py 正是这一设计的直接产物:
_odeint使用@partial(jax.custom_vjp, nondiff_argnums=(0, 1, 2, 3, 4))将func、rtol、atol、mxstep、hmax标记为不可微参数(jax/experimental/ode.py);- 前向
_odeint_fwd返回(ys, (ys, ts, args)),残差正是反向所需的观测点数据(jax/experimental/ode.py); - 反向
_odeint_rev构造增广动力学系统aug_dynamics,用jax.vjp计算伴随方程,并递归调用odeint在反向时间上求解(jax/experimental/ode.py); - 最后
_odeint.defvjp(_odeint_fwd, _odeint_rev)完成绑定(jax/experimental/ode.py)。
这是一个值得精读的完整范例:它展示了
nondiff_argnums、显式残差、以及反向规则内部递归使用自动微分的全部要素。
实践要点小结
- 何时使用 custom_jvp / custom_vjp:当你的函数由底层原语组合而成,但其求导可以给出更数值稳定或更高效的形式(如
logit/expit的稳定梯度);或你的函数需要接入不可微的外部逻辑(求解器、仿真器)而希望指定数学上正确的梯度。 - 记住 JVP 规则中要调用原函数:
defjvp规则体内应调用f本身,这是高阶求导正确性的前提。 - VJP 用显式残差而非闭包:
f_fwd返回(primal, residual),f_bwd消费(residual, cotangent),这一约定让实现简洁且支持vmap场景。 - 不可微参数:用
nondiff_argnums(或nondiff_argnames)标记,注意它们在 JVP 规则中排在primals/tangents之后、在 VJP 反向规则中排在残差之前;不可与defjvps混用。 - 组合性保证:
vmap、jit、grad均可与自定义规则正确组合——vmap会穿过原语作用于底层函数与规则,jit阶段规则已被替换为普通函数调用,无需额外变换规则。 - 限制:对定义了
custom_vjp的函数施加jvp(前向自动微分)不被允许;序列化支持不在范围内。
参考资料
- 设计文档原文:docs/jep/2026-custom-derivatives.md
- 核心实现:jax/_src/custom_derivatives.py
- API 测试:tests/custom_api_test.py
- 规范案例
odeint:jax/experimental/ode.py
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考