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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 08:14:52