JAX 微基准测试怎么测才准:处理 JIT、异步 dispatch 与 32 位 dtype 三个陷阱
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
把一个 NumPy/SciPy 函数移植到 JAX 之后,你通常会想确认它是否真的更快。但直接用%timeit计时得到的数字,在 JAX 上很容易失真:JAX 的运算默认经过 JIT 编译、dispatch 是异步的、默认只使用 32 位 dtype,而 NumPy 默认是 64 位。本文基于 JAX 官方文档 Benchmarking JAX code 与 Asynchronous dispatch,给出一条可照做的微基准测试主路径:在 IPython 中用%time/%timeit分别测出数据传输、JIT 编译和真实运行时,并对齐两边精度,最后按文档标准判断结果是否可信。
先明确 JAX 与 NumPy 计时的四个差异
在动手前需要知道 JAX 计时与 NumPy 的四个根本差异(来自 docs/benchmarking.md):
- JAX 代码是 Just-In-Time(JIT)编译的。即使你自己没有写
jax.jit,JAX 内建函数也是 JIT 编译的,所以第一次运行必然更慢(包含编译)。要拿到 JAX 的最佳性能,应把jax.jit应用到最外层的函数调用上。 - JAX 有异步 dispatch。必须调用
.block_until_ready()才能确保计算真正发生,否则你计时的是“提交任务”的时间而不是“算完”的时间。 - JAX 默认只使用 32 位 dtype。做性能对比时必须对齐精度:64 位计算的成本高于 32 位,所以要么在 NumPy 侧显式使用 32 位 dtype,要么在 JAX 侧启用 64 位 dtype。
- CPU 与加速器之间的数据传输耗时。如果只想测函数求值时间,先把数据传到目标设备,把传输时间单独测出来。
搭建基线微基准:IPython 的 %time 与 %timeit
以下代码来自文档原文,需要在 IPython 环境(交互式 shell 或 notebook)中运行,%time/%timeit是 IPython 的 magic 命令。它同时适用于 NumPy 和 JAX 同一个函数f:
import numpy as np import jax def f(x): # function we're benchmarking (works in both NumPy & JAX) return x.T @ (x - x.mean(axis=0)) x_np = np.ones((1000, 1000), dtype=np.float32) # same as JAX default dtype %timeit f(x_np) # measure NumPy runtime # measure JAX device transfer time %time x_jax = jax.device_put(x_np).block_until_ready() f_jit = jax.jit(f) %time f_jit(x_jax).block_until_ready() # measure JAX compilation time %timeit f_jit(x_jax).block_until_ready() # measure JAX runtime这段代码里每个 magic 各测一件事,不能混读:
%timeit f(x_np):NumPy 侧的运行时。注意输入显式指定dtype=np.float32,与 JAX 的默认 dtype 一致——这就是文档说的“在 NumPy 侧显式用 32 位”的做法。%time x_jax = jax.device_put(x_np).block_until_ready():单独测把数组传上加速器设备的时间。device_put返回的是未完成的 future,.block_until_ready()确保传输真正完成后才停表。- 第一次
%time f_jit(x_jax).block_until_ready():JIT 编译时间。 - 第二次
%timeit f_jit(x_jax).block_until_ready():编译完成后的重复求值时间,这才是“运行时”。
编译与运行被分成两次计时,正是为了隔离“第一次慢”这个 JIT 特征:第一次执行付了编译开销,之后执行才反映真实吞吐。
文档中的示例结果
文档给出在 Colab GPU 上运行上述基准的示例输出(文档示例,你的硬件和版本会得到不同数值,不能作为固定预期):
- NumPy 每次求值 16.2 ms(CPU)
- JAX 把 NumPy 数组拷贝到 GPU 花 1.26 ms
- JAX 编译函数花 193 ms
- JAX 每次求值 485 µs(GPU)
按这个示例,数据传输和函数编译完成后,JAX 在 GPU 上的重复求值约为 NumPy/CPU 的 30 倍。
陷阱二:异步 dispatch 会让 %time 只测到提交时间
这是最隐蔽的失真来源。docs/async_dispatch.rst 解释:JAX 执行jnp.dot(x, x)这类操作时,不会等设备算完再返回,而是返回一个jax.Array——一个“未来才会产生值”的 future。只有当你真正在 host 上读取该值(打印、转成numpy.ndarray)时,JAX 才会强制等待计算完成。
文档给出的 doctest 示例(文档示例,数值不可作为固定预期)展示了这个失真:对一个 1000x1000 的矩阵乘法,
>>> x = random.uniform(random.key(0), (1000, 1000)) >>> %time jnp.dot(x, x) Wall time: 269 µs # 只测到了 dispatch 时间,不是执行时间269 µs 对 CPU 上 1000x1000 的矩阵乘法来说小得可疑——它只是“把任务提交给设备”的时间。要测真实成本,必须二选一(均来自文档):
>>> %time np.asarray(jnp.dot(x, x)) # 在 host 上读取值,会阻塞等待 Wall time: 8.09 ms >>> %time jnp.dot(x, x).block_until_ready() # 不传回 Python,只等计算完成 Wall time: 4.92 ms文档结论:block_until_ready()阻塞但不把结果传回 Python,通常比转回 NumPy 更快,写计算耗时微基准时通常是最优选。这也是上面基线代码里每条%time/%timeit都带.block_until_ready()的原因。
陷阱三:32 位 vs 64 位 dtype 没对齐,比较就不公平
JAX 默认 32 位浮点(在 GPU/TPU 上通常是你要的),NumPy 默认 64 位。两边 dtype 不一致时,你测的其实是精度差异而不是框架差异。docs/101/arrays.md 说明了 JAX 的默认行为:
- 默认情况下 JAX 完全禁用 64 位 dtype:请求
float64会得到float32数组(并给出 warning)。 - 如需 64 位精度,用配置项显式开启:
import jax jax.config.update("jax_enable_x64", True)因此对齐精度有两条等价路径,按文档各选其一即可:
- 主路径(文档默认做法):在 NumPy 侧显式构造 32 位输入,如
np.ones((1000, 1000), dtype=np.float32),与 JAX 默认 dtype 一致——基线代码已经这样做。 - 可选分支:需要严格 64 位对比时,用
jax.config.update("jax_enable_x64", True)打开 JAX 的 64 位,同时 NumPy 侧保持默认 64 位。注意文档提示:64 位计算成本高于 32 位,这本身就是比较的一部分,不要在混精度下下结论。
另外 docs/201/profiling.md 在“Benchmarking JAX code”一节给出的表述与 docs/benchmarking.md 一致:做性能对比时确保两边在相同精度下运行(matched precision)。
怎么判断微基准的结果是否可信
文档对“这是不是公平比较”给了明确的判断标准(见 docs/benchmarking.md 末尾),可以直接用作自检:
- 工作负载是否足够大。文档强调要选足够大的数组(示例是 1000x1000)和足够密集的计算(示例的
@是矩阵-矩阵乘法),才能把 JAX/加速器相对 NumPy/CPU 的额外开销摊薄。 - 反例验证:如果把同样的示例换成 10x10 输入,JAX/GPU 反而比 NumPy/CPU 慢 10 倍(文档示例数值:100 µs vs 10 µs)。小工作负载下加速器的固定开销占主导,此时的“JAX 更慢”不能外推到大负载场景。
- 最终意义在完整应用。文档提醒:真正重要的是完整应用的运行表现,而完整应用不可避免地包含一部分数据传输和编译时间。微基准只回答“算子层面的相对速度”,不要把它当作应用级性能结论。
局限与下一步
微基准回答的是“耗时多少”,不回答“时间花在哪”。如果你已经按上述三步修正后仍想定位开销来源,docs/201/profiling.md 给出了下一级工具:用jax.profiler.start_trace/stop_trace(或jax.profiler.trace上下文管理器)捕获执行 trace,配合 XProf 查看 per-device 时间线;其中 trace 窗口内同样需要用block_until_ready()确保 on-device 执行被完整捕获。此外文档列出的 trace 常见信号——设备时间线空隙(对应逐算子 dispatch、关键路径上的 Python 工作)、异常长的第一步(编译/重 trace)——可以直接指导你回到本节的三个陷阱去修正基准代码。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考