jax.lax.fori_loop上下界相等时仍执行循环体引发索引错误
JAX
jax.lax.fori_loop 上下界均为0时异常触发索引越界的问题 我在代码中使用jax.lax.fori_loop,根据官方文档说明,当设置upper <= lower时应不会产生迭代,直接返回init_val。但当上下界均为0时,循环体相关代码似乎仍被编译检查,进而引发索引越界错误。
复现代码
import jax.numpy as jnp import jax from jax.scipy.special import gammaln # PRELIMINARY PART FOR MWE def comb(n, k): return jnp.round(jnp.exp(gammaln(n + 1) - gammaln(k + 1) - gammaln(n - k + 1))) def binom_conv(n, Aks, Bks): return part_binom_conv(n, 0, n, Aks, Bks) def part_binom_conv(n, k0, k1, Aks, Bks): A_shape = Aks.shape[1:] A_dtype = Aks.dtype init_conv = jnp.zeros(A_shape, dtype=A_dtype) conv = jax.lax.fori_loop(k0, k1, update_binom_conv, (init_conv, n, Aks, Bks))[0] return conv def update_binom_conv(k, val): conv, n, Aks, Bks = val conv = conv + comb(n-1, k) * Aks[k] @ Bks[(n-1)-k] return conv, n, Aks, Bks # IMPORTANT PART def build(U, Hks): n = Hks.shape[0] # n=0 H_shape = Hks.shape[1:] # H_shape=(2,2) Uks_shape = (n+1,)+H_shape # Uks_shape=(1,2,2) Uks = jnp.zeros(Uks_shape, dtype=Hks.dtype) Uks = Uks.at[0].set(U) Uks = jax.lax.fori_loop(0, n, update_Uks, (Uks, Hks))[0] # n=0, so lower=upper=0. Should produce no iterations??? return Uks def update_Uks(k, val): Uks, Hks = val Uks = Uks.at[k+1].set(-1j*binom_conv(k+1, Hks, Uks)) return Uks, Hks # Test Hks = jnp.zeros((0,2,2), dtype=complex) U = jnp.eye(2, dtype=complex) build(U, Hks)
错误信息
--------------------------------------------------------------------------- IndexError Traceback (most recent call last) Cell In[10], line 47 45 Hks = jnp.zeros((0,2,2), dtype=complex) 46 U = jnp.eye(2, dtype=complex) ---> 47 build(U, Hks) Cell In[10], line 35 33 Uks = jnp.zeros(Uks_shape, dtype=Hks.dtype) 34 Uks = Uks.at[0].set(U) ---> 35 Uks = jax.lax.fori_loop(0, n, update_Uks, (Uks, Hks))[0] # n=0, so lower=upper=0. Should produce no iterations??? 36 return Uks [... skipping hidden 12 frame] Cell In[10], line 40 38 def update_Uks(k, val): 39 Uks, Hks = val ---> 40 Uks = Uks.at[k+1].set(-1j*binom_conv(k+1, Hks, Uks)) 41 return Uks, Hks Cell In[10], line 12 11 def binom_conv(n, Aks, Bks): ---> 12 return part_binom_conv(n, 0, n, Aks, Bks) ... --> 930 raise IndexError(f"index is out of bounds for axis {x_axis} with size 0") 931 i = _normalize_index(i, x_shape[x_axis]) if normalize_indices else i 932 i_converted = lax.convert_element_type(i, index_dtype) IndexError: index is out of bounds for axis 0 with size 0
我对此感到困惑,按照文档描述,fori_loop应该直接返回初始值,为何会引发该错误?
问题原因与解决方法
这是JAX即时编译(JIT)的特性导致的:即使fori_loop在运行时不会执行循环体,JAX在编译阶段仍会对循环体函数做静态分析与类型检查,包括验证数组索引的有效性。当传入的Hks是形状为(0,2,2)的数组时,循环体里的binom_conv会尝试访问Aks[k](即Hks[k]),而Hks第0轴长度为0,编译阶段就会触发索引越界错误——哪怕这个代码路径在运行时根本不会被执行。
解决方法
- 提前分支判断:在调用
fori_loop前先判断n > 0,仅满足条件时才执行循环,否则直接返回初始值:def build(U, Hks): n = Hks.shape[0] H_shape = Hks.shape[1:] Uks_shape = (n+1,)+H_shape Uks = jnp.zeros(Uks_shape, dtype=Hks.dtype) Uks = Uks.at[0].set(U) # 提前判断,避免不必要的循环编译检查 if n > 0: Uks = jax.lax.fori_loop(0, n, update_Uks, (Uks, Hks))[0] return Uks - 动态索引保护:在循环体中对数组访问添加边界检查,比如用
jnp.where或jax.lax.dynamic_slice确保索引不会越界,但这种方式会增加运行时开销,不如提前分支高效。
需要注意,JAX的编译逻辑基于静态形状,即使运行时不会触发的代码路径,编译阶段也会进行严格检查,这是它与普通Python代码的核心区别之一。
内容的提问来源于stack exchange,提问作者Ben
相关产品推荐
相关产品推荐

