JAX checkify 变换指南:为 jit / pmap / pjit / scan 代码注入可函数化的运行时错误检查

发布时间:2026/9/10 23:08:35
JAX checkify 变换指南:为 jit / pmap / pjit / scan 代码注入可函数化的运行时错误检查 JAX checkify 变换指南为 jit / pmap / pjit / scan 代码注入可函数化的运行时错误检查【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读本指南以 JAX 官方调试文档 docs/debugging/checkify_guide.md 为主体系统讲解jax.experimental.checkify变换的工作原理与实战用法。你将掌握如何用checkify.check编写 assert 式的运行时断言、如何用checkify.checkify把它们“函数化”为普通返回值从而在jit、pmap、pjit、scan、grad等所有 JAX 变换下自由组合还能学会用float_checks、index_checks等错误集合自动插桩越界索引、NaN 与除零检查并深入理解其底层实现错误值Error、check原语与 discharge 规则最终安全、可靠地保护你的数值计算。checkify 是什么把会抛异常的检查变成返回错误值的函数checkify是 JAX 提供的运行时错误检查变换。它的核心思路是将副作用式的断言像 Python 的assert一样可能抛出异常重写为函数式纯计算——检查失败的谓词被转化为布尔运算错误信息被“排线”plumb成函数的额外输出值最终返回一个携带错误信息的Error值而不是直接中断程序。最简单的用法是把checkify.check当作 assert 用再用checkify.checkify变换包裹目标函数得到一个新的、可以被jax.jit编译的函数from jax.experimental import checkify import jax import jax.numpy as jnp def f(x, i): checkify.check(i 0, index needs to be non-negative, got {i}, ii) y x[i] z jnp.sin(y) return z jittable_f checkify.checkify(f) err, z jax.jit(jittable_f)(jnp.ones((5,)), -2) print(err.get()) # index needs to be non-negative, got -2! (check failed at ...:6 (f))注意三件事返回结构变了checkify之后的函数返回(err, 原输出)二元组其中err是Error类型的值。err.get()只读不抛没有错误时返回None有错误时返回首条错误消息字符串。err.throw()主动抛出有错误时抛出ValueError源码中为JaxRuntimeError它继承自ValueError见 jax/_src/checkify.py无错误时什么都不做。自动插桩常见错误errors 参数除了手动写checkcheckify还能自动为代码插桩常见错误。把多个错误类别用集合运算组合后传给errors参数errors checkify.user_checks | checkify.index_checks | checkify.float_checks checked_f checkify.checkify(f, errorserrors) err, z checked_f(jnp.ones((5,)), 100) err.throw() # ValueError: out-of-bounds indexing at ..:7 (f) err, z checked_f(jnp.ones((5,)), -1) err.throw() # ValueError: index needs to be non-negative! (check failed at …:6 (f)) err, z checked_f(jnp.array([jnp.inf, 1]), 0) err.throw() # ValueError: nan generated by primitive sin at ...:8 (f) err, z checked_f(jnp.array([5, 1]), 0) err.throw() # if no error occurred, throw does nothing!errors的默认值是user_checks即只放行你手写的check不自动插桩。可选的错误类别在源码中有明确定义jax/_src/checkify.py错误集合组成触发条件对应异常类user_checks{FailedCheckError}手写的checkify.check谓词求值为 FalseFailedCheckErrornan_checks{NaNError}浮点运算产生 NaN 输出NaNErrordiv_checks{DivisionByZeroError}发生除零DivisionByZeroErrorfloat_checksnan_checks \| div_checksNaN 或除零上述两者index_checks{OOBError}索引越界OOBErrorautomatic_checksfloat_checks \| index_checks全部自动检查上述三者all_checksautomatic_checks \| user_checks所有检查含手写 check全部这些类别通过frozenset组合因此支持\|集合运算例如checkify.float_checks | checkify.user_checks。更详细的 API 文档可参考 docs/jax.experimental.checkify.rst。为什么 JAX 需要 checkify函数式纯性与 XLA 的约束普通 assert 在部分变换下可用但在 jit 下失效只使用jax.grad和jax.numpy时普通 Python 断言依然有效因为此时数值是即时求值的def f(x): assert x 0., must be positive! return jnp.log(x) jax.grad(f)(0.) # ValueError: must be positive!但一旦进入jit、pmap、pjit、scan计算会被“staged out”延迟为抽象计算图数值在 Python 执行阶段不可见断言立刻崩溃jax.jit(f)(0.) # ConcretizationTypeError: Abstract tracer value encountered ...同样地如果直接jax.jit一个含check但未经过checkify的函数会得到专门设计的报错错误消息常量functionalization_error定义在 jax/_src/checkify.pyjax.jit(f)(jnp.ones((5,)), -1) # checkify transformation not used # ValueError: Cannot abstractly evaluate a checkify.check which was not functionalized.更深层的原因是XLA HLO 不支持断言或抛错即使 JAX 能写出可 staged 的断言 API也没有办法将其 lower 成 XLA。同时JAX 变换的组合语义依赖函数式纯性functional purity一个带隐式抛异常副作用的操作会破坏这种组合性。手工排线plumbing是一种可行但痛苦的做法你可以不依赖任何新 API手动把错误变成返回值在函数外部再抛出def f_checked(x): error x 0. result jnp.log(x) return error, result err, y jax.jit(f_checked)(0.) if err: raise ValueError(must be positive!) # ValueError: must be positive!这个函数是函数式纯的所以天然兼容jit、pmap、pjit、scan以及所有 JAX 变换。问题在于这种排线工作繁琐易错尤其是当检查点很多、错误信息需要携带运行时数值时。checkify 替你做排线checkify正是自动化这条路径它重写你的函数把每次check变成布尔运算、把结果与已跟踪的错误值合并、最终把错误值作为额外输出返回。从源码看这个“重写”发生在checkify的主流程中jax/_src/checkify.py先trace_to_jaxpr把函数跟踪成 jaxpr再调用checkify_jaxpr逐条方程eqn解释执行——每遇到一个原语就查error_checks字典找到对应的检查规则按需插入检查逻辑。这就是文档所说的functionalizing / discharging函数化 / 放电效果。def f(x): checkify.check(x 0., {} must be positive!, x) # convenient but effectful API return jnp.log(x) f_checked checkify(f) err, x jax.jit(f_checked)(-1.) err.throw() # ValueError: -1. must be positive! (check failed at ...:2 (f))checkify.check还支持把运行时值嵌入错误消息通过位置参数*fmt_args和关键字参数**fmt_kwargs提供格式化参数这些参数可以是 traced 值从而在运行时而非跟踪期填充到消息中。需要留意的是跟踪这些数组会带来额外的内存开销——即使最终没有错误发生也是如此源码 docstring 中明确说明见 jax/_src/checkify.py。底层实现Error 值、check 原语与 discharge 规则理解几个源码关键点能帮助你更自信地使用checkify。Error 值是一个 PyTreeError类jax/_src/checkify.py内部维护了每个错误类别的谓词_pred、错误码_code、元数据_metadata和载荷_payload并实现了get()、get_exception()、throw()三个公开方法。它本身注册为 PyTree 节点因此可以像普通值一样被jit、vmap、pmap等变换搬运、切分。这也解释了“错误就是值”errors are just values这一设计哲学。check 是一个带效果的原语checkify.check底层的check_p原语被标记为is_effectful Truejax/_src/checkify.py并且它没有普通的 impl在非debug模式下、未被 discharge 时直接抛出functionalization_error。它注册了各平台的 lowering 规则CPU 和 GPU 上通过emit_python_callback回调到 Python 侧抛错TPU 上则直接标记为不支持jax/_src/checkify.py。discharge 规则与自动插桩规则check_discharge_rulejax/_src/checkify.py负责把check_p的效果“放电”属于enabled_errors的检查被合并进返回的错误值不属于的则被“recharge”重新封存。而自动插桩则通过error_checks字典注册到具体原语上nan_primitives列表jax/_src/checkify.py覆盖了sin、mul、div、sqrt、dot_general、reduce等一大批可能产生 NaN 的 lax 原语nan_error_check会在输出后追加一次jnp.isnan检查index_checks对应的越界检查注册在dynamic_slice、dynamic_update_slice、gather、scatter系列原语上jax/_src/checkify.py例如gather_error_check会计算(start_indices 0) | (start_indices upper_bound)掩码div_error_check同时检查除零与 NaNjax/_src/checkify.py控制流与高阶原语cond、scan、while、pjit、remat、shard_map、custom_jvp_call、custom_vjp_call也都有对应的错误检查规则保证变换可组合。测试用例tests/checkify_test.py中覆盖了这些行为例如test_jit_nan、test_jit_oob、test_jit_div_errors、test_dynamic_slice_oobs、test_cond_basic、test_scan_carry、test_while_loop_body_error等tests/checkify_test.py可作为验证与阅读入口。checkify 与 JAX 变换的组合checkify 之后的函数是函数式纯的因此与所有 JAX 变换“平凡组合”trivially compose。jit先 checkify 再 jit、或先 jit 再 checkify 都可行def f(x, i): return x[i] checkify_of_jit checkify.checkify(jax.jit(f)) jit_of_checkify jax.jit(checkify.checkify(f)) err, _ checkify_of_jit(jnp.ones((5,)), 100) err.get() # out-of-bounds indexing at ..:2 (f) err, _ jit_of_checkify(jnp.ones((5,)), 100) # out-of-bounds indexing at ..:2 (f)vmap / pmap映射map一个 checkify 后的函数会得到一个mapped error——每个映射维度元素可能携带不同的错误def f(x, i): checkify.check(i 0, index needs to be non-negative!) return x[i] checked_f checkify.checkify(f, errorscheckify.all_checks) errs, out jax.vmap(checked_f)(jnp.ones((3, 5)), jnp.array([-1, 2, 100])) errs.throw() ValueError: at mapped index 0: index needs to be non-negative! (check failed at ...:2 (f)) at mapped index 2: out-of-bounds indexing at ...:3 (f) 而 checkify-of-vmap先 vmap 再 checkify只会产生单个未映射的错误——因为 check 是在映射之后的 jaxpr 里放电的jax.vmap def f(x, i): checkify.check(i 0, index needs to be non-negative!) return x[i] checked_f checkify.checkify(f, errorscheckify.all_checks) err, out checked_f(jnp.ones((3, 5)), jnp.array([-1, 2, 100])) err.throw() # ValueError: index needs to be non-negative! (check failed at ...:2 (f))pmap的行为与此类似对 checkified 函数做pmap会得到 mapped errorerr.throw()会逐条列出每个 mapped index 上的失败。pjitpjit包裹 checkified 函数“直接可用”只需给错误值输出额外指定out_axis_resourcesNonedef f(x): return x / x f checkify.checkify(f, errorscheckify.float_checks) f pjit( f, in_shardingsPartitionSpec(x, None), out_shardings(None, PartitionSpec(x, None))) with jax.sharding.Mesh(mesh.devices, mesh.axis_names): err, data f(input_data) err.throw() # ValueError: divided by zero at ...:4 (f)从源码看pjit_error_checkjax/_src/checkify.py会自动为额外的错误值参数补上未指定的UNSPECIFIED入分片、None布局与不捐赠标记并在输出侧同样补上对应的分片因此你只需在用户侧把 error 输出的分片标注为None即全复制。grad前向与反向都会被插桩checkify-of-grad 会让梯度计算同样被自动插桩def f(x): return x / (1 jnp.sqrt(x)) grad_f jax.grad(f) err, _ checkify.checkify(grad_f, errorscheckify.nan_checks)(0.) print(err.get()) # nan generated by primitive mul at ...:3 (f)注意f里没有乘法但它的梯度计算里有乘法——NaN 正是在那里产生的。所以用 checkify-of-grad 可以同时给前向和反向传播运算加自动检查。但checkify.check只会作用在函数输出的primal原始值上如果想检查梯度值需要借助custom_vjp在反向规则里放置checkjax.custom_vjp def assert_gradient_negative(x): return x def fwd(x): return assert_gradient_negative(x), None def bwd(_, grad): checkify.check(grad 0, gradient needs to be negative!) return (grad,) assert_gradient_negative.defvjp(fwd, bwd) jax.grad(assert_gradient_negative)(-1.) # ValueError: gradient needs to be negative!进阶 APIcheck_error 与 debug_checkcheck_error把 Error 值重新变成可函数化的抛错check_error(error)jax/_src/checkify.py语义上等价于err.throw()但它本身可以被checkify函数化因此能放进 staged 代码中。典型场景在函数内部把一段逻辑 checkify 后 jit再在 jit 之外把 Error 值“重新注入”为 Python 异常且整个外层函数仍可再次被 checkifyimport jax from jax.experimental import checkify def f(x): checkify.check(x0, must be positive!) return x def with_inner_jit(x): checked_f checkify.checkify(f) # a checkified function can be jitted error, out jax.jit(checked_f)(x) checkify.check_error(error) return out _ with_inner_jit(1) # no failed check with_inner_jit(-1) # ValueError: must be positive! # can re-checkify the whole thing: error, _ checkify.checkify(with_inner_jit)(-1)简单说checkify把“可函数化的异常效果”变成 Error 值check_error是它的逆操作——把 Error 值变回“可函数化的异常效果”。debug_check只在 checkify 下生效的检查debug_checkjax/_src/checkify.py在未被checkify变换时是 no-op调试标志debugTrue使其在原语 impl 阶段直接跳过只有被checkify变换后才会真正执行——适合放一些开销较大的“调试期断言”def f(x): checkify.debug_check(x!0, cannot be zero!) return x _ f(0) # running without checkify means no debug_check is run. checked_f checkify.checkify(f) err, out jax.jit(checked_f)(0) # running with checkify runs debug_check. err.throw() # ValueError: cannot be zero!另外check要求谓词是标量布尔shape 为()且 dtype 为 bool否则会抛出TypeError见 jax/_src/checkify.py 的is_scalar_pred校验格式化参数必须是数组的 PyTree不允许普通 Python 对象。优势与局限Strengths优势处处可用错误是“普通值”在jit、pmap、pjit、vmap、scan、grad等变换下像其他值一样自然搬运、映射、切分。自动插桩无需逐处修改业务代码checkify可以自动为整段计算加 NaN、除零、越界检查。信息丰富错误消息可携带运行时数值格式化参数并能定位到出错的源文件行号与原始函数名如check failed at ...:6 (f)。Limitations局限检查有开销大量运行时检查代价不菲——例如给每个原语都加 NaN 检查会显著增加计算图中的运算数量。需要显式排线错误错误值必须手动从函数中传出来并主动throw()/get()否则可能悄然错过错误。throw 是阻塞操作抛出错误值会在 host 端物化该值这是阻塞操作会打断 JAX 的异步 run-ahead 执行。相关源码与文档索引官方调试指南docs/debugging/checkify_guide.md公开 API 导出层checkify、check、check_error、debug_check、各错误集合jax/experimental/checkify.py核心实现Error值、check_p原语、discharge 与自动插桩规则、checkify/check/debug_check/check_error入口jax/_src/checkify.py完整测试套件覆盖 jit/pmap/cond/scan/while/pjit/grad 等组合场景tests/checkify_test.pySphinx API 文档docs/jax.experimental.checkify.rst【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考