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

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)
---&gt; 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)
---&gt; 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
---&gt; 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):
---&gt; 12     return part_binom_conv(n, 0, n, Aks, Bks)
...
--&gt; 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,编译阶段就会触发索引越界错误——哪怕这个代码路径在运行时根本不会被执行。

解决方法

  1. 提前分支判断:在调用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
    
  2. 动态索引保护:在循环体中对数组访问添加边界检查,比如用jnp.where或jax.lax.dynamic_slice确保索引不会越界,但这种方式会增加运行时开销,不如提前分支高效。

需要注意,JAX的编译逻辑基于静态形状,即使运行时不会触发的代码路径,编译阶段也会进行严格检查,这是它与普通Python代码的核心区别之一。


内容的提问来源于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 06:14:53