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

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

2026-09-09 11:20:45作者:胡唯隽

如果你的 JAX 程序已经被 jax.jitjax.vmapjax.grad 包裹,你仍然需要在运行时执行一段带副作用的 host 端 Python 代码(写文件、读取外部状态、更新全局变量等),普通的 print 或 Python 函数调用就不够用了:在编译后的计算图里,它们看到的是 trace 期的抽象值,而不是运行时数据。jax.experimental.io_callback 就是为这类场景设计的回调原语:它在运行时把数据从设备传回 host 进程,执行你传入的 Python 函数,再把结果送回计算中。本文基于 JAX 官方教程 External callbacks(另一份等价教程见 docs/external-callbacks.md)梳理它的基本用法,以及它在 jitvmapgradscan/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_callbackvmap 兼容的前提是 ordered=Falsevmapscan/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.pyio_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 都支持

scanwhile_loopio_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.pyio_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_dtypesNone 表示回调无输出;运行时 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 同硬件),这部分传输通常是快速零拷贝的;
  • vmapordered=False 时回调顺序不保证;
  • 相关的回归测试位于 tests/python_callback_test.py,可用于对照本文各行为描述。

更多 jax.debug.print / jax.debug.callback 的调试细节见 docs/debugging/print_breakpoint.md

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
docsdocs
暂无描述
Markdown
900
5.83 K
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
927
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.94 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
603
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
396
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.04 K
527