news 2026/9/10 10:32:55

深入理解 JAX 自定义导数规则:custom_jvp / custom_vjp 设计原理与实现剖析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
深入理解 JAX 自定义导数规则:custom_jvp / custom_vjp 设计原理与实现剖析

深入理解 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_jvpjax.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 中,有两种方式可以定义求导规则:

  1. 使用jax.custom_jvpjax.custom_vjp,为已经是 JAX 可变换(JAX-transformable)的 Python 函数定义自定义的前向(JVP)与反向(VJP)求导规则。这是本文的绝对主题。
  2. 定义新的core.Primitive实例并为其实现全部变换规则,例如接入求解器、仿真器等外部数值系统的函数调用。这是更高阶、更低层的方式,本文不展开。

作为 JAX 开发者,我们希望以logitexpit(见 jax/scipy/special.py 中的相关实现)为代表,用其他原语定义库函数,但在求导时表现出"原语级"行为——即显式给出可能更数值稳定性能更好的自定义导数规则,同时不必为这些函数单独指定vmapjit规则。作为长远目标(stretch goal),还希望让高阶函数(fixed_pointodeintroot等)能方便地接入自定义求导规则。

设计文档明确指出本设计要解决的问题清单:关闭 #116、#1097、#1249、#1275 等一系列 issue,并取代旧的custom_transforms机制。


设计目标(Goals)与非目标(Non-goals)

核心目标

设计文档将目标明确划分为两级:

  • 用户侧:希望用户能自定义其代码的前向/反向求导行为,且该自定义
    1. 具备清晰一致的语义,能正确与其他 JAX 变换组合;
    2. 足够灵活,支持 Autograd、PyTorch 中常见的用法,包括对 Python 控制流求导、NaN 调试等场景。
  • 开发者侧logit/expit这类以原语组合定义的库函数,求导时具有"原语级"自定义规则,而无需额外提供vmap/jit规则。

由此归纳出的主要目标是:

  1. 解决vmap 移除自定义 JVP 语义问题(对应 issue #1249);
  2. 允许在自定义 VJP 中使用 Python(例如调试 NaN,对应 issue #1275)。

次要目标包括:简化用户体验(符号零、kwargs 等);推动用户能轻松为fixed_pointodeintroot等添加自定义规则。

明确排除的非目标

  1. 不做变换泛化的自定义机制:旧的custom_transforms目标是"变换泛化"地自定义行为,理论上允许用户为任意变换定制规则。本设计只解决求导(JVP 与 VJP 分开)的自定义问题——这是实际唯一被请求的场景;通过专精化降低了复杂度、提升了灵活性。若要控制所有规则,直接编写原语即可。
  2. 不优先追求数学美学:虽然自定义 VJP 的签名a -> (b, CT b --o CT a)在数学上更优雅,但 Python 机制难以处理返回类型中的闭包,因此实现上**显式处理残差(residuals)**而非依赖闭包。
  3. 暂不支持序列化:将 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.] ← 意外!

最后一行gradvmap的结果不符合预期。一般来说,施加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_primitivevmap透明的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)的同时进行的。这带来两个关键好处:

  1. 支持基于 pdb 的工作流:用户可以用标准 Python 调试器检查数值、捕获 NaN。设计文档作者特别提到,在实现odeint原语期间多次依靠运行时数值检查来调试问题;一个特别实用的技巧是在自定义 VJP 规则中插入调试器断点,从而在反向传播的特定位置进入调试器。
  2. 允许对 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_callcore.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原语,并预先将ff_jvp追踪为 jaxpr,以简化 initial-style 处理。

脚注:虽然道德上应在绑定custom_jvp_call_jaxpr前就为ff_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)

  1. 任意 pytree:输入与输出类型abc可以是任意 jaxtypes 的 pytree(嵌套的 tuple/list/dict 等)。
  2. 按名传参(kwargs):当 kwargs 能通过inspect模块解析为位置参数时支持按名传参。设计文档称其为一次"对 Python 3 增强的程序化签名检查能力的实验",可靠但不完备
  3. nondiff_argnums标记不可微参数:与jitstatic_argnums类似,这些参数不要求是 JAX 类型。其传递约定为:设主函数签名为(d, a) -> bd为不可微类型),则
    • JVP 规则签名为(a, T a, d) -> T b——不可微参数按顺序排在primalstangents之后
    • VJP 规则的反向函数签名为(d, c, CT b) -> CT a——不可微参数按顺序排在残差之前

当前仓库源码中的 API 佐证

仓库实现 jax/_src/custom_derivatives.py 与设计文档完全对应:

  • CustomJVPCallPrimitivecustom_jvp_call_p)与CustomVJPCallPrimitivecustom_vjp_call_p)分别对应设计中的两个 Python 级调用原语(见 jax/_src/custom_derivatives.py 与 jax/_src/custom_derivatives.py)。
  • CustomJVP__init__中,nondiff_argnames会通过fun_signatureinfer_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_partialprepend_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

仓库实现印证了设计文档中的几个关键实现决策:

  1. 每个变换都有自定义绑定方法:为custom_jvp_callcustom_vjp_call提供了类似core.call_bind的自定义 bind 方法,区别在于不处理 env traces(那些直接报错)。

  2. 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_linvmap与 MLIR lowering 都直接抛错(raise_custom_vjp_error_on_jvp),即对自定义 VJP 函数施加jvp是不被允许的——这与设计文档"施加 jvp 时若定义了自定义 VJP 规则会失败"的表述一致。此外,对custom_vjp函数施加前向自动微分会触发disallow_jvp(jax/_src/custom_derivatives.py)。

  3. 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))funcrtolatolmxstephmax标记为不可微参数(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、显式残差、以及反向规则内部递归使用自动微分的全部要素。


实践要点小结

  1. 何时使用 custom_jvp / custom_vjp:当你的函数由底层原语组合而成,但其求导可以给出更数值稳定或更高效的形式(如logit/expit的稳定梯度);或你的函数需要接入不可微的外部逻辑(求解器、仿真器)而希望指定数学上正确的梯度。
  2. 记住 JVP 规则中要调用原函数defjvp规则体内应调用f本身,这是高阶求导正确性的前提。
  3. VJP 用显式残差而非闭包f_fwd返回(primal, residual)f_bwd消费(residual, cotangent),这一约定让实现简洁且支持vmap场景。
  4. 不可微参数:用nondiff_argnums(或nondiff_argnames)标记,注意它们在 JVP 规则中排在primals/tangents之后、在 VJP 反向规则中排在残差之前;不可与defjvps混用。
  5. 组合性保证vmapjitgrad均可与自定义规则正确组合——vmap会穿过原语作用于底层函数与规则,jit阶段规则已被替换为普通函数调用,无需额外变换规则。
  6. 限制:对定义了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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 10:32:33

4 步解锁 WeMod Pro:Wand-Enhancer 从克隆到生效的完整路径

4 步解锁 WeMod Pro:Wand-Enhancer 从克隆到生效的完整路径 【免费下载链接】Wand-Enhancer Advanced UX and interoperability extension for Wand (WeMod) app 项目地址: https://gitcode.com/GitHub_Trending/we/Wand-Enhancer 按下面的步骤走完&#xff…

作者头像 李华
网站建设 2026/9/10 10:30:37

TVBoxOSC:给电视盒子装一个每天自动更新的 TVBox 播放器

TVBoxOSC:给电视盒子装一个每天自动更新的 TVBox 播放器 【免费下载链接】TVBoxOSC TVBoxOSC - 一个基于第三方项目的代码库,用于电视盒子的控制和管理。 项目地址: https://gitcode.com/GitHub_Trending/tv/TVBoxOSC TVBoxOSC 是一个面向 Androi…

作者头像 李华
网站建设 2026/9/10 10:29:21

Flutter跨平台化学学习应用开发与鸿蒙适配实践

1. 项目背景与核心价值 化学元素与方程式学习一直是初高中理科教育的重点难点。传统纸质教材存在互动性差、记忆效率低等问题,而市面上多数学习类APP又往往功能单一、平台受限。我们团队基于Flutter框架开发的这款跨平台应用,首次实现了在鸿蒙系统上的原…

作者头像 李华