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

如何将含for循环的JAX代码转换为jax.lax.scan以加速编译?

问题:JAX中用jax.lax.scan替代带jit的for循环时出现静态索引错误

原可运行代码(带jax.jit的for循环)

import functools
import jax
import jax.numpy as jnp

@functools.partial(jax.jit, static_argnums=0)
def func(n):
    p = 1
    x = jnp.arange(8)
    y = jnp.zeros((n,))

    for idx in range(n):
        y = y.at[idx].set(jnp.sum(x[::p]))
        p = 2*p

    return y

func(2)
# >> Array([28., 12.], dtype=float32)

错误的scan转换代码(报静态索引错误)

import numpy as np

def body(p, xi):
    y = jnp.sum(x[::p])
    p = 2*p
    return p, y

x = jnp.arange(8)
jax.lax.scan(body, 1, np.arange(2))
# >> IndexError: Array slice indices must have static start/stop/step ...

错误原因

JAX的切片操作x[::p]要求步长p是静态编译时常量。原jit版本能运行,是因为n被标记为静态参数,循环次数固定,每次迭代的p值在编译时可被推导为静态常量;但在jax.lax.scan中,p作为循环的动态状态变量,编译时无法确定其值,因此直接用动态p做切片步长会触发错误。

正确的scan转换方法

我们可以用动态掩码筛选元素替代切片操作,避免静态步长限制。具体来说,通过判断索引是否是p的倍数来筛选需要求和的元素:

import jax
import jax.numpy as jnp

x = jnp.arange(8)
x_indices = jnp.arange(len(x))  # 预生成x的索引数组

def body(p, _):
    # 筛选索引是p倍数的元素并求和
    sum_val = jnp.sum(x[x_indices % p == 0])
    new_p = 2 * p
    return new_p, sum_val

# 用scan执行循环,输入的xs可以是任意长度为n的数组(这里用jnp.arange(n)占位)
n = 2
final_p, y = jax.lax.scan(body, init=1, xs=jnp.arange(n))

print(y)
# >> Array([28., 12.], dtype=float32)

如果需要将整个逻辑包装成jit函数,也可以把n设为静态参数(和原逻辑一致):

import functools

@functools.partial(jax.jit, static_argnums=0)
def func_scan(n):
    x = jnp.arange(8)
    x_indices = jnp.arange(len(x))
    
    def body(p, _):
        sum_val = jnp.sum(x[x_indices % p == 0])
        return 2*p, sum_val
    
    _, y = jax.lax.scan(body, 1, xs=jnp.arange(n))
    return y

print(func_scan(2))
# >> Array([28., 12.], dtype=float32)

补充优化

如果p的取值规律是2的幂次,还可以通过位运算优化掩码判断(当p是2的幂时,x_indices & (p-1) == 0等价于模p等于0),性能会略有提升:

sum_val = jnp.sum(x[(x_indices & (p - 1)) == 0])

内容的提问来源于stack exchange,提问作者DavidJ

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 23:27:08