资讯动态

如何在 jit、vmap、grad 内部用 jax.io_callback 执行 host 端 Python 代码

发布时间:2026/9/10 12:33:45 来源:尧图企业网站定制
如何在 jit、vmap、grad 内部用 jax.io_callback 执行 host 端 Python 代码【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax如果你的 JAX 程序已经被jax.jit、jax.vmap或jax.grad包裹你仍然需要在运行时执行一段带副作用的 host 端 Python 代码写文件、读取外部状态、更新全局变量等普通的print或 Python 函数调用就不够用了在编译后的计算图里它们看到的是 trace 期的抽象值而不是运行时数据。jax.experimental.io_callback就是为这类场景设计的回调原语它在运行时把数据从设备传回 host 进程执行你传入的 Python 函数再把结果送回计算中。本文基于 JAX 官方教程 External callbacks另一份等价教程见 docs/external-callbacks.md梳理它的基本用法以及它在jit、vmap、grad、scan/while_loop下的行为边界。为什么普通 print 不行以及 io_callback 的定位先看一个典型的踩坑现象在jit内部用print打印中间变量打印出来的不是运行时数值而是 trace 期的抽象值import jax jax.jit def f(x): y x 1 print(intermediate value: {}.format(y)) return y * 2 result f(2)要在运行时拿到真实值需要走回调机制。如果只是打印调试输出可以直接用jax.debug.printjax.jit def f(x): y x 1 jax.debug.print(intermediate value: {}, y) return y * 2 result f(2)其原理是把y的运行时值作为 CPUjax.Array传回 host 进程执行打印。JAX 提供三种回调旧版唯一的jax.experimental.host_callback已弃用按场景选用jax.pure_callback适合纯函数无副作用可能被编译器省略或重复调用jax.experimental.io_callback适合不纯函数读/写磁盘、更新全局状态等这是本文主题jax.debug.callback适合需要严格反映编译器执行行为的函数不能返回任何值。三者与转换的兼容性对照来自教程中的表格callback functionsupports return valuejitvmapgradscan/with while_loopguaranteed executionjax.pure_callback支持支持支持不支持可配合custom_jvp支持不保证jax.experimental.io_callback支持支持仅orderedFalse时支持不支持支持保证执行jax.debug.callback不支持支持支持支持支持不保证需要注意两点脚注io_callback与vmap兼容的前提是orderedFalsevmap套scan/while_loop再套io_callback的语义较复杂文档明确说明其行为可能在未来版本变化。在 jit 内部调用 io_callbackio_callback的函数签名为见 jax/_src/callback.py 中的实现def io_callback( callback, # 在 host 上执行的 Python 函数假定带副作用 result_shape_dtypes, # 描述回调输出的 pytree叶节点需有 shape/dtype 属性 *args, # 传给 callback 的实参 shardingNone, # 可选指定从哪个设备发起回调 orderedFalse, # 是否要求顺序调用 **kwargs, ): ...文档给出的示例是调用一个全局 host 端 NumPy 随机数生成器。这是一个典型的不纯操作打印是副作用global_rng的状态更新也是副作用文档同时说明这只是一个演示例子并非在 JAX 中生成随机数的推荐方式import jax import jax.numpy as jnp import numpy as np from jax.experimental import io_callback global_rng np.random.default_rng(0) def host_side_random_like(x): Generate a random array like x using the global_rng state # We have two side-effects here: # - printing the shape and dtype # - calling global_rng, thus updating its state print(fgenerating {x.dtype}{list(x.shape)}) return global_rng.uniform(sizex.shape).astype(x.dtype) jax.jit def numpy_random_like(x): return io_callback(host_side_random_like, x, x) x jnp.zeros(5) numpy_random_like(x)第二个参数result_shape_dtypes在这里直接传入了与输出形状/类型一致的x实际使用中它应是结构匹配回调输出的 pytree常用jax.ShapeDtypeStruct定义叶节点。验证方式运行后 host 进程会打印一行类似generating float32[5]的输出文档示例输出实际 dtype 取决于运行环境同时global_rng的状态被更新——再次运行会消费到不同的随机数。这正是副作用真的发生了的判据。与pure_callback的一个关键区别即使回调的输出在后续计算中没有被使用编译器也不会删掉io_callback的执行。在 vmap 内部默认支持但执行顺序不保证io_callback默认orderedFalse可以直接被vmapjax.vmap(numpy_random_like)(x)但要记住mapped 的各次回调可能以任意顺序执行。文档指出在 GPU 上运行时各 mapped 输出的顺序可能逐次运行都不同。如果你的逻辑依赖回调顺序例如依赖全局状态的连续更新设置orderedTrue。此时对结果做vmap会直接报错源码中的报错信息为见 jax/_src/callback.py 中io_callback_batching_ruleValueError: Cannot vmap ordered IO callback.jax.jit def numpy_random_like_ordered(x): return io_callback(host_side_random_like, x, x, orderedTrue) jax.vmap(numpy_random_like_ordered)(x) # 抛出上面的 ValueError这个报错本身就是一个明确的验证点如果你预期顺序但不想 vmap出现该异常说明配置按预期生效。在 scan / while_loop 内部无论是否 ordered 都支持scan和while_loop与io_callback组合时不受ordered标志影响def body_fun(_, x): return _, numpy_random_like_ordered(x) jax.lax.scan(body_fun, None, jnp.arange(5.0))[1]即使用orderedTrue的版本放进scan也能正常工作。文档同时提醒vmapofscan/while_loopofio_callback的语义复杂行为可能在未来发布中变化涉及这一组合时不要依赖当前的具体行为。在 grad 内部不能依赖被求导的变量io_callback没有自动求导规则。源码中 JVP 和 transpose 规则均直接抛出ValueError(IO callbacks do not support JVP.)见 jax/_src/callback.py 中io_callback_jvp_rule。因此在grad下如果回调依赖被求导的变量该变量会作为实参传入回调求导会失败例如对上一节的numpy_random_like求导会抛异常如果回调不依赖任何被求导变量则可以正常执行例如jax.jit def f(x): io_callback(lambda: print(hello), None) return x jax.grad(f)(1.0);这里result_shape_dtypes传None表示回调无输出运行时 host 进程会打印hello文档示例输出。第二个参数不需要与任何返回值对应因为该回调只产生副作用。分片环境下的行为可选分支如果程序运行在多设备分片环境回调运行在 host、编译计算之外它在哪里跑、看到什么取决于模式详见教程 Callbacks and sharding 一节全局视图模式explicit 或 auto sharding参数会被 gather 到单个设备回调在该设备的 host 上执行一次拿到完整的全球值。语义一致但大规模下 gather 可能变慢或超出内存jax.shard_map的 full manual 模式没有全局视图回调按设备逐个执行每次只看到自己的分片——这是分片局部日志、按 host 加载数据的典型模式。限制与下一步不要把io_callback用于求导路径上依赖被求导变量的计算需要可导的 host 端函数时文档给出的路径是jax.pure_callback配合jax.custom_jvp手动定义求导规则教程中有一个完整的 Bessel 函数scipy.special.jv包装示例可在 docs/201/callbacks.md 查看每次回调都会触发设备到 host 的数据传输与同步。在 GPU/TPU 等加速器上这是明显的开销在单 CPU 上host 与 device 同硬件这部分传输通常是快速零拷贝的vmap下orderedFalse时回调顺序不保证相关的回归测试位于 tests/python_callback_test.py可用于对照本文各行为描述。更多jax.debug.print/jax.debug.callback的调试细节见 docs/debugging/print_breakpoint.md。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价