You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为何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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 07:08:16