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

如何在JAX的jit装饰函数中实现带掩码的索引取值?

JAX中JIT模式下通过掩码索引取值的解决方法

错误原因

添加@jax.jit后触发NonConcreteBooleanIndexError,是因为JIT编译要求数组形状在编译阶段确定。直接用布尔索引过滤coords会得到长度动态变化的数组(取决于掩码中True的数量),而原代码中使用的x.at[...].get()高级索引方式无法处理这种动态形状,因此报错。

解决方案

改用jnp.gather_nd函数,它原生支持动态长度的索引数组,能在JIT模式下正常工作。具体实现如下:

import jax.numpy as jnp
from jax import jit

x = jnp.arange(25).reshape((5,5))
coords = jnp.array([
    [1,2],
    [2,3],
    [1,2],
    [1,2],
])
coords_mask = jnp.array([True, True, False, True])

@jit
def masked_gather(x, coords, coords_mask):
    # 过滤出掩码为True的坐标
    coords_masked = jnp.boolean_mask(coords, coords_mask)
    # 使用gather_nd直接获取对应位置的值
    return jnp.gather_nd(x, coords_masked)

# 执行结果:Array([ 7, 13,  7], dtype=int32)
print(masked_gather(x, coords, coords_mask))

另一种实现方式(通过整数索引)

也可以先将布尔掩码转换为整数索引,再提取坐标并取值,效果一致:

@jit
def masked_gather(x, coords, coords_mask):
    # 获取掩码为True的位置索引
    valid_indices = jnp.nonzero(coords_mask)[0]
    # 提取有效坐标
    coords_masked = jnp.take(coords, valid_indices, axis=0)
    return jnp.gather_nd(x, coords_masked)

这两种方法都能在JIT编译下正常运行,返回预期结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 19:50:03