如何在 jit、vmap、grad 内部用 jax.io_callback 执行 host 端 Python 代码
如果你的 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.print:
@jax.jit
def f(x):
y = x + 1
jax.debug.print("intermediate value: {}", y)
return y * 2
result = f(2)
其原理是把 y 的运行时值作为 CPU jax.Array 传回 host 进程执行打印。
JAX 提供三种回调,旧版唯一的 jax.experimental.host_callback 已弃用,按场景选用:
jax.pure_callback:适合纯函数(无副作用),可能被编译器省略或重复调用;jax.experimental.io_callback:适合不纯函数(读/写磁盘、更新全局状态等),这是本文主题;jax.debug.callback:适合需要严格反映编译器执行行为的函数,不能返回任何值。
三者与转换的兼容性对照(来自教程中的表格):
| callback function | supports return value | jit |
vmap |
grad |
scan/with while_loop |
guaranteed execution |
|---|---|---|---|---|---|---|
jax.pure_callback |
支持 | 支持 | 支持 | 不支持(可配合 custom_jvp) |
支持 | 不保证 |
jax.experimental.io_callback |
支持 | 支持 | 仅 ordered=False 时支持 |
不支持 | 支持 | 保证执行 |
jax.debug.callback |
不支持 | 支持 | 支持 | 支持 | 支持 | 不保证 |
需要注意两点脚注:io_callback 与 vmap 兼容的前提是 ordered=False;vmap 套 scan/while_loop 再套 io_callback 的语义较复杂,文档明确说明其行为可能在未来版本变化。
在 jit 内部调用 io_callback
io_callback 的函数签名为(见 jax/_src/callback.py 中的实现):
def io_callback(
callback, # 在 host 上执行的 Python 函数,假定带副作用
result_shape_dtypes, # 描述回调输出的 pytree,叶节点需有 shape/dtype 属性
*args, # 传给 callback 的实参
sharding=None, # 可选,指定从哪个设备发起回调
ordered=False, # 是否要求顺序调用
**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(f'generating {x.dtype}{list(x.shape)}')
return global_rng.uniform(size=x.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 默认(ordered=False)可以直接被 vmap:
jax.vmap(numpy_random_like)(x)
但要记住:mapped 的各次回调可能以任意顺序执行。文档指出,在 GPU 上运行时,各 mapped 输出的顺序可能逐次运行都不同。
如果你的逻辑依赖回调顺序(例如依赖全局状态的连续更新),设置 ordered=True。此时对结果做 vmap 会直接报错,源码中的报错信息为(见 jax/_src/callback.py 中 io_callback_batching_rule):
ValueError: Cannot `vmap` ordered IO callback.
@jax.jit
def numpy_random_like_ordered(x):
return io_callback(host_side_random_like, x, x, ordered=True)
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]
即使用 ordered=True 的版本放进 scan 也能正常工作。文档同时提醒:vmap of scan/while_loop of io_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下ordered=False时回调顺序不保证;- 相关的回归测试位于 tests/python_callback_test.py,可用于对照本文各行为描述。
更多 jax.debug.print / jax.debug.callback 的调试细节见 docs/debugging/print_breakpoint.md。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0634
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
jforgamejforgame是一个一站式游戏服务器开发框架。包含游戏服务器开发所需要的各种组件,比如网关,socket服务端与客户端,自定义高效消息编解码,游戏热更新,游戏通用工具等等。包含游戏服,跨服,匹配服,后台管理系统等实现,同时提供大量业务案例以供学习。亦可用于其他socket应用,例如及时聊天等。Java01
fizz-gateway-nodeAn Aggregation API Gateway in Java . FizzGate 是一个基于 Java开发的微服务聚合网关,是拥有自主知识产权的应用网关国产化替代方案,能够实现热服务编排聚合、自动授权选择、线上服务脚本编码、在线测试、高性能路由、API审核管理、回调管理等目的,拥有强大的自定义插件系统可以自行扩展,并且提供友好的图形化配置界面,能够快速帮助企业进行API服务治理、减少中间层胶水代码以及降低编码投入、提高 API 服务的稳定性和安全性。Java00
certd开源SSL证书管理工具;全自动证书申请、更新、续期;通配符证书,泛域名证书申请;证书自动化部署到阿里云、腾讯云、主机、群晖、宝塔;https证书,pfx证书,der证书,TLS证书,nginx证书自动续签自动部署JavaScript00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00