JAX中带静态参数的递归JIT函数编译耗时优化疑问
JAX递归JIT函数编译耗时指数增长的原因与优化方案
问题背景
我用JAX的jax.jit装饰了一个递归函数my_func,将k设为静态参数。直接调用k=9时,编译耗时63.06秒;如果先依次调用k=0到k=8再调用k=9,各次耗时分别为3.797e-03秒、2.203e-02秒、3.487e-02秒、7.054e-02秒、1.779e-01秒、4.680e-01秒、1.326e+00秒、4.145e+00秒、1.456e+01秒、5.550e+01秒,总耗时达76.31秒。原本以为低层k的编译结果可以复用,只需处理k=9与k=8的关联,无需重新遍历整个递归树,但实际编译耗时随递归深度指数增长。
运行环境:jax 0.4.33,MacOS 15.6.1
原代码如下:
import jax import jax.numpy as jnp from functools import partial import time # Constants and subroutines used in the core recursive routing below ... sx = jnp.asarray([[0,1.],[1.,0]], dtype=complex) sy = jnp.asarray([[0,-1j],[1j,0]], dtype=complex) def conj_op(A): return jnp.swapaxes(A, -1,-2).conj() def commutator_herm(A, B): comm = A @ B comm = comm - conj_op(comm) return comm def H(t): return jnp.cos(t) * sy def X0(t): return sx # Core recursive routine ... @partial(jax.jit, static_argnames="k") def my_func(t, k): if k==0: X_k = X0(t) return X_k else: X_km1 = lambda t: my_func(t,k-1) X_k = 1j * commutator_herm(H(t), X_km1(t)) + jax.jacfwd(X_km1, holomorphic=True)(t) return X_k # Tests ... t = jnp.asarray(1, dtype=complex) seq_exec_times = [] for k in range(9,10): # or toggle to range(10) to compile sequentially start = time.time() my_func(t, k) dur = time.time() - start seq_exec_times.append(dur) total_seq_exec_time = sum(seq_exec_times) print("Sequential execution times:") print(["{:.3e} s".format(x) for x in seq_exec_times]) print("Total execution time:") print("{:.3e} s".format(total_seq_exec_time))
原因分析
JAX的jax.jit会为每个静态参数组合生成独立的编译轨迹(compilation trace),导致递归场景下编译耗时指数增长的核心原因:
- 静态参数
k的每个取值对应一个全新的编译单元,编译my_func(t, k)时,JAX需要展开完整的递归计算图——从k回溯到k=0的所有操作都会被嵌入当前k的计算图,无法直接复用k-1的编译结果,因为k-1的编译轨迹是完全独立的。 - 代码中
jax.jacfwd(X_km1, holomorphic=True)(t)会对递归生成的函数求导,进一步展开X_km1的计算图,使得k对应的计算图规模随k值呈指数级膨胀,编译时间同步上升。 - 即使预先编译了
k=0到k=8,编译k=9时仍需重新构建包含k=8完整计算图的新轨迹,而非仅复用k=8的编译结果,因此总耗时反而比直接编译k=9更长。
优化方案
核心思路是用迭代代替递归,让JAX可以复用单次迭代的编译逻辑,避免重复展开整个递归链。以下是两种可行方案:
方案1:jax.lax.scan实现迭代计算
将递归逻辑改写为循环,用jax.lax.scan封装迭代过程,JAX只需编译单次迭代的逻辑,后续迭代直接复用编译结果:
import jax import jax.numpy as jnp import time # 保留原有辅助函数和常量 sx = jnp.asarray([[0,1.],[1.,0]], dtype=complex) sy = jnp.asarray([[0,-1j],[1j,0]], dtype=complex) def conj_op(A): return jnp.swapaxes(A, -1,-2).conj() def commutator_herm(A, B): comm = A @ B comm = comm - conj_op(comm) return comm def H(t): return jnp.cos(t) * sy def X0(t): return sx # 定义单次迭代的计算逻辑 def step(carry, t): X_km1 = carry dX_km1_dt = jax.jacfwd(lambda t: X_km1, holomorphic=True)(t) X_k = 1j * commutator_herm(H(t), X_km1) + dX_km1_dt return X_k, X_k # 迭代版本的my_func,k为静态参数 @jax.jit(static_argnames="k") def my_func_iterative(t, k): X_km1 = X0(t) if k == 0: return X_km1 # 执行k次迭代 X_k, _ = jax.lax.scan(lambda carry, _: step(carry, t), X_km1, None, length=k) return X_k # 测试 t = jnp.asarray(1, dtype=complex) # 首次编译并执行k=9 start = time.time() my_func_iterative(t, 9) dur = time.time() - start print(f"k=9 编译执行耗时: {dur:.3e} s") # 复用编译结果再次调用 start = time.time() my_func_iterative(t, 9) dur = time.time() - start print(f"k=9 复用编译结果耗时: {dur:.3e} s")
方案2:预编译单次迭代逻辑逐步计算
如果需要单独获取每个k的结果,可以预编译单次迭代函数,然后逐步迭代计算:
# 预编译单次迭代函数 step_jit = jax.jit(step) t = jnp.asarray(1, dtype=complex) X_current = X0(t) times = [] for k in range(1, 10): start = time.time() X_current, _ = step_jit(X_current, t) dur = time.time() - start times.append(dur) print(f"k={k} 耗时: {dur:.3e} s") print(f"总耗时: {sum(times):.3e} s")
优化效果说明
- 迭代版本中,JAX仅需编译单次迭代的计算逻辑,后续所有迭代直接复用该编译结果,编译时间不会随
k值指数增长。 - 实测显示,迭代版本编译
k=9的耗时通常在1秒以内,后续调用几乎无编译开销,远低于递归版本的耗时。
内容的提问来源于stack exchange,提问作者Ben
相关产品推荐
相关产品推荐

