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

