ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

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

深入理解 JAX 自定义导数规则:custom_jvp / custom_vjp 设计原理与实现剖析 深入理解 JAX 自定义导数规则custom_jvp / custom_vjp 设计原理与实现剖析【免费下载链接】jaxComposable transformations of PythonNumPy 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相当于内联inliningf_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 相同但不同于 TF1JAX 对 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_jvpjvp(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-stylejaxpr 形成原语是个例外它对 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) - bd为不可微类型则JVP 规则签名为(a, T a, d) - T b——不可微参数按顺序排在primals与tangents之后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_signature与infer_argnums_and_argnames解析为nondiff_argnumsjax/_src/custom_derivatives.pydefjvp支持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_vmaptests/custom_api_test.py分别验证vmap(f)、vmap(jvp(f))、jvp(vmap(f))、vmap(jvp(vmap(f)))的结果都与期望一致——这正是设计文档vmap 穿过 custom_jvp_call语义的直接测试。nondiff_argnums/nondiff_argnamestests/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_jvpjax/_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.pyAPI 测试tests/custom_api_test.py规范案例odeintjax/experimental/ode.py【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表