JAX是否保存JIT编译函数的jaxpr?形状变化时编译机制解析
示例代码
import jax import jax.numpy as jnp @jax.jit def test(x): if x.shape[0] > 4: return 1 else: return -1 print(test(jnp.ones(8,))) print(test(jnp.ones(3,)))
运行结果
1 -1
你观察到的现象背后是JAX的JIT缓存与静态分支折叠机制在起作用,下面逐个解答你的疑问:
1. JAX是否保存JIT编译函数的jaxpr?
是的。JAX会将JIT编译生成的jaxpr以及对应的XLA可执行代码缓存起来,存储在jax.jit装饰器维护的内部缓存结构中,避免对相同输入签名的函数重复编译。
2. jaxpr是否每次调用唯一?
不是。jaxpr的唯一性由输入的静态参数决定——包括输入数组的形状、dtype,以及函数中通过static_argnums/static_argnames标记的静态参数。只要两次调用的静态参数完全一致,就会复用同一个jaxpr;静态参数不同时,会生成新的jaxpr。
3. 是否仅形状匹配时复用?
不止形状,还需匹配dtype等其他静态属性。比如同样是(8,)形状的数组,float32和float64类型会触发不同的编译流程,生成不同的jaxpr。只有当所有静态参数(形状、dtype等)完全匹配时,才会复用已缓存的jaxpr和编译结果。
4. 形状不同时是否重新编译?
是的。输入形状不同属于静态参数变更,JAX会判定为新的函数签名,重新追踪函数执行过程,生成对应的jaxpr并编译为新的可执行代码。这也是你传入(3,)数组时得到-1的原因:此时JAX重新追踪,发现x.shape[0] = 3,静态判定条件为False,会直接折叠掉if分支,生成仅返回-1的jaxpr,这个jaxpr和第一次的完全独立。
补充说明
你看到首次调用的jaxpr没有分支,是因为JAX在JIT追踪时会折叠静态已知的条件分支:当传入x.shape[0] = 8时,x.shape[0] > 4是静态可判定的True,JAX会直接剔除else分支,只保留返回1的逻辑;传入x.shape[0] = 3时则相反,折叠掉if分支,只保留返回-1的逻辑。这两个是独立的缓存条目,因此会输出不同结果。
内容的提问来源于stack exchange,提问作者Quasi

