ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

JAX 异步派发(Asynchronous Dispatch)深入解析:future 语义、阻塞时机与基准测试正确姿势

JAX 异步派发(Asynchronous Dispatch)深入解析:future 语义、阻塞时机与基准测试正确姿势 JAX 异步派发Asynchronous Dispatch深入解析future 语义、阻塞时机与基准测试正确姿势【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读JAX 通过异步派发Asynchronous Dispatch将 Python 代码与加速器GPU/TPU/CPU的执行解耦让 Python 程序跑在设备前面从而隐藏单次操作的调度开销。本文基于 docs/async_dispatch.rst 展开结合仓库源码jax/_src/array.py、jax/_src/dispatch.py、jaxlib/py_array.cc逐层剖析jax.Array的 future 语义、宿主侧何时会被阻塞以及如何用block_until_ready()写出度量真实计算耗时的微基准测试。一、什么是异步派发用 future 隐藏 Python 开销JAX 之所以能在不牺牲易用性的前提下获得接近原生加速器的性能关键设计之一就是异步派发asynchronous dispatch。考虑下面这个程序import numpy as np import jax.numpy as jnp from jax import random x random.uniform(random.key(0), (1000, 1000)) # 打印结果即执行 repr(result) 或 str(result)会阻塞直到该值就绪 jnp.dot(x, x) 3.当执行jnp.dot(x, x)这类操作时JAX并不会等待操作完成再返回控制权给 Python 程序。相反JAX 立即返回一个 jax.Array 值——它是一个future一个未来才会在加速器设备上被生产出来、但当下不一定可用的值。我们可以在不等待产生它的计算完成的前提下检查jax.Array的 shape 或 dtype甚至可以把它继续传给另一个 JAX 计算例如上面代码里的加法运算。只有当我们在宿主host侧真正检查数组的值时——比如打印它或把它转换成普通的numpy.ndarray——JAX 才会强制 Python 代码等待计算完成。从源码可以看到这一语义的实现结构在 jax/_src/array.py 中ArrayImpl.block_until_ready()遍历其持有的所有底层缓冲区self._arrays并逐一调用db.block_until_ready()而底层的 C 实现在 jaxlib/py_array.cc通过self.BlockUntilReady()等待 XLA 执行结果就绪。也就是说值与计算完成是两个可以分离的概念——前者随手可得后者需要显式同步。二、异步派发的价值让 Python 远离关键路径异步派发之所以有用是因为它允许 Python 代码**跑在加速器前面run ahead**把 Python 从关键路径critical path上移开。只要满足两个前提Python 代码入队enqueue工作的速度快于设备执行这些工作的速度Python 代码实际上不需要在宿主侧检查某个计算的输出那么一个 Python 程序就可以入队任意数量的工作而避免让加速器空等。这正是数据并行、流水线等场景中实现计算与调度重叠的基础设备端忙于执行上一批 kernelPython 端同时在为下一批 kernel 做 tracing 与派发。这一机制在仓库中也有直接的工程体现jax/_src/dispatch.py 中的RuntimeTokenSet维护了每个设备上最后一次派发计算的输出运行时 tokenoutput_runtime_tokens用于在具有副作用的计算如有序副作用 effect之间建立依赖而其block_until_ready()会等待所有 token 就绪并清空集合。更值得注意的是 jax/_src/dispatch.py 中通过atexit.register注册的wait_for_tokens()即使你的程序没有显式同步JAX 也会在 Python 进程退出时自动等待所有已派发的计算完成保证异步派发不会造成程序退出但设备还在跑的悬空状态。这也解释了为什么在 Jupyter 等交互环境中异步派发后的结果最终总能被安全地读取。三、微基准测试的陷阱269µs 的 1000×1000 矩阵乘法异步派发有一个略显意外slightly surprising的后果它在微基准测试中尤其明显 %time jnp.dot(x, x) CPU times: user 267 µs, sys: 93 µs, total: 360 µs Wall time: 269 µs269µs 对一次 CPU 上的 1000×1000 矩阵乘法来说小得惊人然而真相是异步派发误导了我们我们测量的并不是矩阵乘法的执行时间而仅仅是派发dispatch这步工作所花的时间。%time在语句返回时即停止计时而jnp.dot(x, x)在设备真正算完之前就已经返回了那个 future。要测量操作的真实成本必须二选一在宿主侧读取该值例如把它转换为普通的宿主侧 numpy 数组或者调用 jax.Array.block_until_ready() 方法等待产生它的计算完成。四、正确的测量方式np.asarray 与 block_until_ready4.1 方式一转为 numpy 数组阻塞且带回数据 %time np.asarray(jnp.dot(x, x)) CPU times: user 61.1 ms, sys: 0 ns, total: 61.1 ms Wall time: 8.09 ms从源码看np.asarray(jax.Array)最终会触发 jax/_src/array.py 中的_value属性它会为每个分片调用_copy_single_device_array_to_host_async()发起到宿主的异步拷贝随后在_single_device_array_to_np_array_did_copy()中等待数据真正拷贝完成再把结果组装成只读的numpy.ndarray并缓存到_npy_value。因此这种方式既阻塞等待计算完成又把数据搬回了宿主。4.2 方式二block_until_ready阻塞但数据留在设备上 %time jnp.dot(x, x).block_until_ready() CPU times: user 50.3 ms, sys: 928 µs, total: 51.2 ms Wall time: 4.92 msblock_until_ready()的实现jax/_src/array.py只是等待底层每个 buffer 的计算完成不涉及把数据从设备传输回宿主use_cpp_method() def block_until_ready(self): self._check_if_deleted() for db in self._arrays: db.block_until_ready() return self4.3 两者对比与最佳实践阻塞但不把结果传回 Python 通常更快因此block_until_ready()往往是编写测量计算耗时的微基准测试时的最佳选择。它规避了两个额外开销从设备到宿主的数据传输H2D/D2H 拷贝时间在宿主侧组装、缓存 numpy 数组的时间。当你的目标是只测计算本身有多快时用block_until_ready()当你的目标是模拟真实用户读取结果例如整个端到端 pipeline 的延迟时才需要np.asarray()。仓库中还有两个相关的实用 API 值得了解jax.block_until_ready(x)见 jax/_src/api.py针对整个pytree的便捷版本tries to call ablock_until_readymethod on pytree leaves返回结构与输入相同、所有 JAX 数组叶子都已就绪的 pytree。jax.effects_barrier()见 jax/_src/api.py等待所有已有函数完成副作用内部调用dispatch.runtime_tokens.block_until_ready()适用于对有序副作用有同步需求的场景。此外block_until_ready并非jax.Array的专利仓库中的PRNGKeyArrayjax/_src/random/prng.py、EArrayjax/_src/earray.py、core.Tokenjax/_src/core.py以及稀疏数组基类jax/experimental/sparse/_base.py都实现了同名的就绪等待方法形成了统一的显式同步约定。五、测试用例中的佐证block_until_ready 的实际用法在仓库测试中block_until_ready是验证异步执行语义与资源生命周期的常用工具。例如 tests/api_test.py 用f(1, 2).block_until_ready()确保可执行对象确实被执行后再检查client.live_executables()的存活数量tests/api_test.py 则在触发删除错误后调用result.block_until_ready()并断言抛出RuntimeError因为底层 buffer 已被回收_check_if_deleted()会拦截。这提示了一个使用细节block_until_ready()调用前若数组已被显式delete()或垃圾回收会触发RuntimeError见 jax/_src/array.py 的_check_if_deleted。因此在长期运行的 JAX 进程中如果希望异步派发的计算一定完成最好的习惯是尽早、显式地调用block_until_ready()而不是依赖进程退出时的atexit兜底。六、小结操作是否阻塞等待计算是否把数据传回宿主适用场景jnp.dot(x, x)直接返回否立即返回 future否流水线式入队隐藏调度开销print(arr)/repr(arr)/str(arr)是是部分调试、交互式查看np.asarray(jnp.dot(x, x))是是需要 numpy 结果 / 端到端延迟测量jnp.dot(x, x).block_until_ready()是否纯计算耗时微基准测试推荐理解异步派发是正确使用 JAX 性能工具的基石它让 Python 端可以持续超前入队工作把调度成本从关键路径上剥离同时也要求你在做微基准测试时显式同步优先block_until_ready()否则测出来的只是派发时间而非真实计算时间。这一机制贯穿 jax/_src/dispatch.py派发与 token 管理、jax/_src/array.pyjax.Array的 future 语义与阻塞实现以及 jaxlib/py_array.ccC 层的同步原语是理解 JAX 执行模型不可绕过的一环。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表