JAX 外部回调(External Callbacks)完全指南:pure_callback、io_callback 与 debug.callback
2026/9/10 2:20:51 网站建设 项目流程

JAX 外部回调(External Callbacks)完全指南:pure_callback、io_callback 与 debug.callback

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

本教程系统讲解 JAX 中的三类外部回调机制——jax.pure_callbackjax.experimental.io_callbackjax.debug.callback,它们允许 JAX 运行时在**主机端(host)**执行 Python 代码,且可安全用于jax.jitjax.vmapjax.grad等变换之中。读完本文,你将掌握如何在 JIT 编译下打印运行时值、在变换中调用 NumPy/SciPy 等非 JAX 库函数,并通过custom_jvp为回调补齐自动微分规则。

为什么需要回调(Why callbacks)

回调(callback)是一种在运行时于主机侧执行代码的机制。以"在计算过程中打印某个变量的值"为例,直接用 Python 的print在 JIT 函数中打印,得到的并不是运行时的真实值,而是追踪期(trace-time)的抽象值

import jax @jax.jit def f(x): y = x + 1 print("intermediate value: {}".format(y)) # 打印的是抽象值,而非运行时值 return y * 2 result = f(2)

要打印运行时的值,需要借助回调。例如使用jax.debug.print(关于追踪与调试的更多背景,可参阅 key-concepts.md 与 debugging.md):

@jax.jit def f(x): y = x + 1 jax.debug.print("intermediate value: {}", y) return y * 2 result = f(2)

其工作原理是:把y的运行时值作为 CPU 上的jax.Array传回主机进程,主机进程再将其打印出来。这就是外部回调最朴素的形态——跨设备边界把数据送回主机执行

回调的种类(Flavors of callback)

早期版本的 JAX 只有一种回调实现,即jax.experimental.host_callback。该机制存在一些缺陷,现已废弃,取而代之的是面向不同场景设计的三种回调:

  • jax.pure_callback:适用于纯函数(无副作用,如不打印、不读写磁盘、不更新全局状态)。
  • jax.experimental.io_callback:适用于非纯函数(有副作用,例如读写磁盘数据)。
  • jax.debug.callback:适用于需要如实反映编译器执行行为的函数,是通用调试场景的首选。

(前文用到的jax.debug.print本质上是jax.debug.callback的封装。)

从用户视角看,这三种回调的核心区别在于:它们各自允许哪些变换与编译器优化。下表是官方文档给出的完整对照:

回调函数支持返回值jitvmapgradscan/while_loop保证执行
jax.pure_callback❌¹
jax.experimental.io_callback✅/❌²✅³
jax.debug.callback

¹jax.pure_callback可通过custom_jvp与自动微分兼容(见下文示例)。
²jax.experimental.io_callback仅在ordered=False时才与vmap兼容。
³ 注意:对io_callback进行vmap后再套scan/while_loop语义较复杂,其行为可能在后续版本中变化。

深入pure_callback

当你需要在主机侧执行一个纯函数(无副作用)时,jax.pure_callback是首选。传入的函数实际未必是纯的,但 JAX 的变换和高阶函数会假定它是纯的——这意味着它可能被编译器静默消除(elide),也可能被多次调用。

基本用法如下:在回调内调用 NumPy(而非jax.numpy)运算,并用jax.ShapeDtypeStruct声明结果的 shape 与 dtype:

import jax import jax.numpy as jnp import numpy as np def f_host(x): # 调用 numpy(非 jax.numpy)操作: return np.sin(x).astype(x.dtype) def f(x): result_shape = jax.ShapeDtypeStruct.like(x) return jax.pure_callback(f_host, result_shape, x, vmap_method='sequential') x = jnp.arange(5.0) f(x)

由于pure_callback可以被消除或复制,它开箱即用地兼容jit以及scanwhile_loop等高阶原语:

jax.jit(f)(x) def body_fun(_, x): return _, f(x) jax.lax.scan(body_fun, None, jnp.arange(5.0))[1]

因为调用时显式指定了vmap_method,它同样兼容vmap

jax.vmap(f)(x)

然而,由于 JAX 无法内省回调内容,pure_callback没有定义自动微分语义

jax.grad(f)(x) # 报错:pure callbacks do not support JVP

这对应源码中的实现:在 jax/_src/callback.py 中,pure_callback_jvp_rulepure_callback_transpose_rule会直接抛出ValueError(提示改用custom_jvp/custom_vjp)。结合custom_jvp使用pure_callback的完整示例见下文。

vmap_method的取值语义

从 pure_callback 的源码文档 可以看到,vmap_method控制回调在vmap下的变换行为,合法取值为["sequential", "sequential_unrolled", "expand_dims", "broadcast_all", "legacy_vectorized", None],传入非法值会直接抛出ValueError。各取值含义如下:

  • "sequential":使用jax.lax.map沿批量轴循环,对每个 batch 元素调用一次回调。
  • "sequential_unrolled":与sequential类似,但循环被展开(unrolled)。
  • "expand_dims":对未批量化的输入在头部添加大小为 1 的新轴后调用回调。
  • "broadcast_all":与expand_dims类似,但会把输入平铺(tile)成预期的批量 shape。scipy.special.jv这类能原生处理广播输入的库函数适合用此方法。
  • 默认行为说明:当前未显式指定时默认使用"sequential",但该默认行为已被弃用——未来版本默认会改为在未指定时抛出NotImplementedError,因此建议总是显式传入vmap_method

纯函数的消除语义与异常边界

由于设计上假定函数无副作用,若回调的输出未被使用,编译器可能将整个回调消除

def print_something(): print('printing something') return np.int32(0) @jax.jit def f1(): return jax.pure_callback(print_something, np.int32(0)) # 输出被使用,回调执行 f1(); @jax.jit def f2(): jax.pure_callback(print_something, np.int32(0)) # 输出未使用,回调被消除 return 1.0 f2();

f1中回调的输出用于函数返回值,因此回调被执行并打印;而在f2中输出未被使用,编译器发现后直接消除了调用——这正是"无副作用函数"应有的正确语义。

pure_callback与异常:在 JAX 变换的语境下,Python 运行时异常应被视为副作用。因此在pure_callback故意抛错违反 API 契约,程序行为是未定义的:程序如何终止通常取决于后端,且细节可能在后续版本中变化。此外,把非纯函数传给pure_callback,在jit/vmap等变换下可能产生意外行为,因为变换规则建立在"回调是纯的"这一假设之上。例如:

import jax import jax.numpy as jnp def raise_via_callback(x): def _raise(x): raise ValueError(f"value of x is {x}") return jax.pure_callback(_raise, x, x) def raise_if_negative(x): return jax.lax.cond(x < 0, raise_via_callback, lambda x: x, x) x_batch = jnp.arange(4) [raise_if_negative(x) for x in x_batch] # 不抛出异常 jax.vmap(raise_if_negative)(x_batch) # ValueError: value of x is 0

同样一段逻辑,逐元素调用不报错,vmap后却报错。为避免此类问题,官方建议不要试图用pure_callback来抛运行时错误

深入io_callback

pure_callback相反,jax.experimental.io_callback明确面向非纯函数(有副作用)。以下示例回调到主机端的全局 NumPy 随机数生成器——这是非纯操作,因为生成随机数会更新随机状态(注意:这仅是演示io_callback的玩具示例,并非 JAX 推荐的随机数生成方式):

from jax.experimental import io_callback from functools import partial import numpy as np global_rng = np.random.default_rng(0) def host_side_random_like(x): """使用 global_rng 状态生成与 x 同形状的随机数组""" # 这里有两个副作用: # - 打印 shape 和 dtype # - 调用 global_rng,从而更新其状态 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)

io_callback默认兼容vmap

jax.vmap(numpy_random_like)(x)

但要注意:映射后的回调可能以任意顺序执行。例如在 GPU 上运行时,映射输出的顺序可能每次运行都不同。

如果回调的执行顺序很重要,可以设置ordered=True;此时再尝试vmap会报错:

@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) # 报错:Cannot `vmap` ordered IO callback

这一限制在源码中有明确体现:io_callback_batching_ruleordered=True时直接抛出ValueError(见 jax/_src/callback.py)。另一方面,scanwhile_loop无论是否强制排序都可以与io_callback配合:

def body_fun(_, x): return _, numpy_random_like_ordered(x) jax.lax.scan(body_fun, None, jnp.arange(5.0))[1]

pure_callback一样,若io_callback接收了被微分的变量,在自动微分下会失败:

jax.grad(numpy_random_like)(x) # 报错:IO callbacks do not support JVP

但如果回调不依赖被微分的变量,它仍然可以执行:

@jax.jit def f(x): io_callback(lambda: print('hello'), None) return x jax.grad(f)(1.0) # 打印 hello,正常工作

pure_callback不同,即使回调的输出在后续计算中未被使用,编译器也不会移除io_callback的执行(这正是"有副作用"语义的体现,与 io_callback 的实现 中将其标记为带IOEffect/OrderedIOEffect副作用一致)。

深入debug.callback

pure_callbackio_callback都对其调用的函数施加了纯度假设,并在不同程度上限制了 JAX 变换与编译机制。而debug.callback对回调函数几乎不做任何假设——回调的行为如实反映 JAX 在程序执行过程中的实际动作;同时,debug.callback不能向程序返回任何值

from jax import debug def log_value(x): # 这里可以是真正的日志调用;此处用 print() 演示 print("log:", x) @jax.jit def f(x): debug.callback(log_value, x) return x f(1.0)

debug.callback兼容vmap

x = jnp.arange(5.0) jax.vmap(f)(x)

也兼容grad及其他自动微分变换:

jax.grad(f)(1.0)

从源码实现看(jax/_src/debugging.py),debug_callback_p被注册了debug_callback_jvp_rule(返回空切线)与debug_callback_transpose_rule(返回None占位),因此它在grad下安全;同时它被标记为带DebugEffect副作用,且注册了 CPU/GPU/TPU 三个平台的 lowering 规则。正是这种"不假设、只如实反映"的特性,使debug.callback在通用调试场景中比另外两种回调更有用。

示例:pure_callback结合custom_jvp

pure_callbackjax.custom_jvp结合,是一种强大的用法(custom_jvp的更多细节可参阅 advanced_autodiff.md)。

假设你想为某个尚未被jax.scipyjax.numpy包装的 SciPy/NumPy 函数创建 JAX 兼容包装器。这里以第一类贝塞尔函数scipy.special.jv为例。首先定义一个直接的pure_callback

import jax import jax.numpy as jnp import scipy.special def jv(v, z): v, z = jnp.asarray(v), jnp.asarray(z) # 要求阶数 v 为整数类型:这会简化下面的 JVP 规则 assert jnp.issubdtype(v.dtype, jnp.integer) # 将输入提升为非精确类型(float/complex)。 # 注意 jnp.result_type() 会考虑 enable_x64 标志。 z = z.astype(jnp.result_type(float, z.dtype)) # 包装 scipy 函数以返回预期的 dtype _scipy_jv = lambda v, z: scipy.special.jv(v, z).astype(z.dtype) # 定义输出的预期 shape 与 dtype result_shape_dtype = jax.ShapeDtypeStruct( shape=jnp.broadcast_shapes(v.shape, z.shape), dtype=z.dtype) # 使用 vmap_method="broadcast_all",因为 scipy.special.jv 能处理广播输入 return jax.pure_callback(_scipy_jv, result_shape_dtype, v, z, vmap_method="broadcast_all")

这样就能从被变换的 JAX 代码(包括jitvmap变换)中调用scipy.special.jv

from functools import partial j1 = partial(jv, 1) z = jnp.arange(5.0) print(j1(z)) print(jax.jit(j1)(z)) # jit 下的结果 print(jax.vmap(j1)(z)) # vmap 下的结果

但直接调用jax.grad会报错,因为该函数没有定义自动微分规则:

jax.grad(j1)(z) # 报错

接下来为它定义自定义梯度规则。根据第一类贝塞尔函数的定义,关于参数z的导数存在一个简洁的递推关系:

$$ d J_\nu(z) = \left{ \begin{eqnarray} -J_1(z),\ &\nu=0\ [J_{\nu - 1}(z) - J_{\nu + 1}(z)]/2,\ &\nu\ne 0 \end{eqnarray}\right. $$

关于 $\nu$ 的梯度更复杂,但本例中已把v参数限制为整数类型,因此不必为它求导。用jax.custom_jvp定义回调函数的自动微分规则:

jv = jax.custom_jvp(jv) @jv.defjvp def _jv_jvp(primals, tangents): v, z = primals _, z_dot = tangents # 注意:v_dot 恒为 0,因为 v 是整数 jv_minus_1, jv_plus_1 = jv(v - 1, z), jv(v + 1, z) djv_dz = jnp.where(v == 0, -jv_plus_1, 0.5 * (jv_minus_1 - jv_plus_1)) return jv(v, z), z_dot * djv_dz

现在计算梯度就能正确工作了:

j1 = partial(jv, 1) print(jax.grad(j1)(2.0))

更进一步,由于梯度是用jv自身定义的,JAX 的架构意味着二阶及更高阶导数会自动生效

jax.hessian(j1)(2.0)

性能注意事项

虽然以上方案在 JAX 中完全正确,但要注意:每次调用基于回调的jv函数,都会把输入数据从设备传到主机,再把scipy.special.jv的输出从主机传回设备

  • 在 GPU/TPU 等加速器上运行时,这种数据搬运与主机同步会带来显著的开销,每次调用jv都如此。
  • 如果 JAX 运行在单个 CPU 上("主机"与"设备"在同一硬件上),JAX 通常能以零拷贝的快速方式完成数据传输,这使得该模式成为扩展 JAX 能力的一种相对直接的方式。

总结与选型建议

外部回调是连接 JAX 计算图与主机 Python 生态的桥梁,其核心取舍在于"纯度假设"与"变换自由度":

  • 纯函数、需要返回值jax.pure_callback,并显式指定vmap_method;需要求导时配合custom_jvp(如 jax/_src/callback.py 所示,源码中pure_callback的 JVP/transpose 规则默认抛错,必须自定义)。
  • 有副作用、需要保证执行与返回值jax.experimental.io_callback;需要保序时用ordered=True(代价是放弃vmap)。
  • 调试、无返回值、需要如实反映编译器行为jax.debug.callbackjax.debug.print,它们是唯一兼容grad的回调家族。

上述三种回调的语义差异在测试集中也有大量覆盖,例如 tests/python_callback_test.py 内含数十个相关测试用例,可供深入理解各回调在jitvmapgradscan等变换下的预期行为。

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询