如何将含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
相关产品推荐
相关产品推荐

