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

使用jax.lax.scan扫描Equinox模型时出现静态切片索引错误

JAX scan触发动态切片索引错误的原因及解决方法

问题背景

你构建的Equinox模型如下:

import jax
import jax.numpy as jnp, jax.random as jrnd
import equinox as eqx

class Model(eqx.Module):
    lags: list[int]
    linear: eqx.Module

    def __init__(self, lags: list[int]=[22, 5, 1], *, key: jrnd.PRNGKeyArray):
        self.lags = lags
        self.linear = eqx.nn.Linear(len(lags), 1, key=key)

    def __call__(self, x, key=None):
        x_new = jnp.array([
            x[-lag:].mean() for lag in self.lags
        ])
        return self.linear(x_new)

单独调用eqx.filter_jit(loss)或vmap时一切正常,但执行jax.lax.scan(scanner, model, jnp.arange(5))时触发错误:

IndexError: Array slice indices must have static start/stop/step to be used with NumPy indexing syntax. Found slice(Traced<ShapedArray(int32[], weak_type=True)>with<DynamicJaxprTrace(level=2/0)>, None, None). To index a statically sized array at a dynamic position, try lax.dynamic_slice/dynamic_update_slice (JAX does not support dynamically sized arrays within JIT compiled functions).

你疑惑明明self.lags是常量列表,为何只有scan会触发这个错误。

原因分析

self.lags确实是常量,但**jax.lax.scan会把整个model对象当作可追踪的carry状态,默认会追踪model内的所有属性**,包括看起来静态的list。而eqx.filter_jit是Equinox专门针对Module设计的,它能自动识别Module里的静态参数(比如你的lags),不会将其纳入JIT追踪范围,所以单独用filter_jit或vmap时不会有问题。

但jax.lax.scan没有Equinox的静态参数识别逻辑,会把lags当成动态值处理。此时x[-lag:]里的lag就变成了JAX追踪的动态变量,而JAX的NumPy风格切片(x[a:b])要求索引必须是编译期就能确定的静态值,动态索引自然会触发错误。

解决方法

方法1:用eqx.filter_scan替代jax.lax.scan

Equinox的eqx.filter_scan和filter_jit逻辑一致,能自动区分Module中的静态与动态参数,直接替换即可解决问题:

eqx.filter_scan(scanner, model, jnp.arange(5))

方法2:手动标记lags为静态字段

如果一定要用jax.lax.scan,可以给lags加上eqx.static_field()标记,明确告诉JAX这个参数是静态不变的:
修改Model类的定义:

class Model(eqx.Module):
    lags: list[int] = eqx.static_field()  # 标记为静态字段
    linear: eqx.Module

    def __init__(self, lags: list[int]=[22, 5, 1], *, key: jrnd.PRNGKeyArray):
        self.lags = lags
        self.linear = eqx.nn.Linear(len(lags), 1, key=key)

这样即使在jax.lax.scan中,lags也会被当作静态值处理,切片索引不会被追踪成动态变量。

方法3:改用JAX动态切片API(通用方案)

如果遇到确实需要动态索引的场景(你的场景不需要,但可作为参考),可以用jax.lax.dynamic_slice替代NumPy风格切片:
将x[-lag:].mean()替换为:

start_idx = x.shape[0] - lag
slice_x = jax.lax.dynamic_slice(x, (start_idx,), (lag,))
slice_x.mean()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 09:23:15