为何JAX的JIT编译在我的测试示例中第二次运行更慢?
问题分析与解决
为什么第二次运行反而更慢?
主要有两个核心原因:
1. 超大规模数据引发虚拟内存交换
你的测试用例中,x是形状为(1_000_000, 1_000)的numpy数组,按默认float64类型计算,内存占用约8GB(1e6 * 1e3 * 8字节)。如果你的机器物理内存不足,第一次运行时系统会被迫将部分数据写入磁盘交换区;第二次运行时,需要把交换区的数据重新加载回内存,磁盘IO的巨大开销直接导致运行时间变长。
2. JIT与vmap的组合方式不合理
你当前的代码是先对单个my_function做JIT,再用vmap包装它。这种方式下,vmap会对每个输入批次单独调用JIT后的函数,无法让JAX将整个向量循环编译为单一的设备内核,调度开销会抵消JIT的优化效果,甚至引发额外性能损耗。
修正后的代码示例
from icecream import ic import jax from time import time import numpy as np def my_function(x, y): return x @ y # 关键修正:对vmap后的整体函数做JIT,让JAX编译整个向量操作 vectorized_function = jax.jit(jax.vmap(my_function, in_axes=(0, None))) # 改用适合物理内存的测试形状,避免虚拟内存干扰 shape = (100_000, 1_000) x = np.ones(shape) y = np.ones(shape[1]) # 预热调用:单独触发编译,避免编译时间影响计时结果 vectorized_function(x, y) # 正式计时 start = time() vectorized_function(x, y) t_1 = time() - start start = time() vectorized_function(x, y) t_2 = time() - start print(f'{t_1 = }\n{t_2 = }')
额外说明
- 预热调用:第一次JIT调用会触发编译,单独执行这一步可以让后续计时只反映实际计算开销。
- 数据规模控制:确保测试数据能完全放入物理内存,排除磁盘交换的干扰,才能准确观察JIT的缓存效果。
- JIT包装顺序:始终优先对
vmap/pmap等批量操作后的整体函数做JIT,才能最大化设备内核的编译优化。
内容的提问来源于stack exchange,提问作者kernel123
相关产品推荐
相关产品推荐

