尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

JAX custom_vjp 迁移指南:nondiff_argnums 限制更新与闭包修复详解

JAX custom_vjp 迁移指南:nondiff_argnums 限制更新与闭包修复详解 机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载导读本文基于 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与输出余切cotangentg返回一个长度等于原始函数参数个数的元组代表各参数的梯度。装饰器还接受可选的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_argnumsnondiff_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_argnumslo、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若结构不符会抛出带详细说明的TypeErrorbwd 只接收两个参数(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 isnt 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 cant 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_gradientif 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_hashableFalse) 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 #4039JEP 文档还预告了 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)二元组需要保留下来的非可微数组必须存入 residualbwd 规则签名改为(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 源码完成存量代码的安全迁移。赞分享机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载相关推荐JAX custom_vjp 与 nondiff_argnums 升级指南从闭包式非可微参数迁移到残差式自定义 VJPJAX 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数据分析上一篇Blender 2.8→3.6 项目升级避坑指南从材质到Python的全流程迁移方案下一篇内存安全检测技术演进从ASAN原始论文到现代硬件加速方案终极指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表