拓冰建站拓冰建站
首页 / 资讯中心 / 正文

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

JAX 外部回调External Callbacks完全指南pure_callback、io_callback 与 debug.callback【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本教程系统讲解 JAX 中的三类外部回调机制——jax.pure_callback、jax.experimental.io_callback与jax.debug.callback它们允许 JAX 运行时在**主机端host**执行 Python 代码且可安全用于jax.jit、jax.vmap、jax.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.mdjax.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仅在orderedFalse时才与vmap兼容。³ 注意对io_callback进行vmap后再套scan/while_loop语义较复杂其行为可能在后续版本中变化。深入pure_callback当你需要在主机侧执行一个纯函数无副作用时jax.pure_callback是首选。传入的函数实际未必是纯的但 JAX 的变换和高阶函数会假定它是纯的——这意味着它可能被编译器静默消除elide也可能被多次调用。基本用法如下在回调内调用 NumPy而非jax.numpy运算并用jax.ShapeDtypeStruct声明结果的 shape 与 dtypeimport 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_methodsequential) x jnp.arange(5.0) f(x)由于pure_callback可以被消除或复制它开箱即用地兼容jit以及scan、while_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它同样兼容vmapjax.vmap(f)(x)然而由于 JAX 无法内省回调内容pure_callback没有定义自动微分语义jax.grad(f)(x) # 报错pure callbacks do not support JVP这对应源码中的实现在 jax/_src/callback.py 中pure_callback_jvp_rule与pure_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(fvalue 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(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)io_callback默认兼容vmapjax.vmap(numpy_random_like)(x)但要注意映射后的回调可能以任意顺序执行。例如在 GPU 上运行时映射输出的顺序可能每次运行都不同。如果回调的执行顺序很重要可以设置orderedTrue此时再尝试vmap会报错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) # 报错Cannot vmap ordered IO callback这一限制在源码中有明确体现io_callback_batching_rule在orderedTrue时直接抛出ValueError见 jax/_src/callback.py。另一方面scan和while_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.callbackpure_callback与io_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兼容vmapx jnp.arange(5.0) jax.vmap(f)(x)也兼容grad及其他自动微分变换jax.grad(f)(1.0)从源码实现看jax/_src/debugging.pydebug_callback_p被注册了debug_callback_jvp_rule返回空切线与debug_callback_transpose_rule返回None占位因此它在grad下安全同时它被标记为带DebugEffect副作用且注册了 CPU/GPU/TPU 三个平台的 lowering 规则。正是这种不假设、只如实反映的特性使debug.callback在通用调试场景中比另外两种回调更有用。示例pure_callback结合custom_jvp将pure_callback与jax.custom_jvp结合是一种强大的用法custom_jvp的更多细节可参阅 advanced_autodiff.md。假设你想为某个尚未被jax.scipy或jax.numpy包装的 SciPy/NumPy 函数创建 JAX 兼容包装器。这里以第一类贝塞尔函数scipy.special.jv为例。首先定义一个直接的pure_callbackimport 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( shapejnp.broadcast_shapes(v.shape, z.shape), dtypez.dtype) # 使用 vmap_methodbroadcast_all因为 scipy.special.jv 能处理广播输入 return jax.pure_callback(_scipy_jv, result_shape_dtype, v, z, vmap_methodbroadcast_all)这样就能从被变换的 JAX 代码包括jit与vmap变换中调用scipy.special.jvfrom 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),\ \nu0\ [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需要保序时用orderedTrue代价是放弃vmap。调试、无返回值、需要如实反映编译器行为→jax.debug.callback或jax.debug.print它们是唯一兼容grad的回调家族。上述三种回调的语义差异在测试集中也有大量覆盖例如 tests/python_callback_test.py 内含数十个相关测试用例可供深入理解各回调在jit、vmap、grad、scan等变换下的预期行为。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门