如何在JIT函数中使用jnp.where获取最后一个索引
JAX JIT环境下获取符合阈值条件的最后一个元素
问题场景
现有两个JAX数组:
z = jnp.array([[5.55751118], [5.18212974], [4.35981727], [3.4559711 ], [3.35750248], [2.65199945], [2.02298999], [1.59444971], [0.80865185], [0.77579791]]) z1 = jnp.array([[ 1.58559484], [ 3.79094097], [-0.52712522], [-1.0178286 ], [-3.51076985], [ 1.30108161], [-1.29824303], [-0.19209007], [ 0.37451138], [-2.33619987]])
需求是从z的第一行开始,找到z1中对应z >= (z[0]-2.6)条件的最后一个值。非JIT环境下可通过jnp.where实现,但添加@jit装饰器后会抛出ConcretizationTypeError——原因是jnp.where返回的索引数组长度是动态的,JIT无法静态确定其大小,导致追踪值无法被具体化。
解决方案
改用静态形状的操作替代jnp.where,通过反转掩码数组找到第一个符合条件的位置,再转换为原数组的索引:
import jax.numpy as jnp from jax import jit z = jnp.array([[5.55751118], [5.18212974], [4.35981727], [3.4559711 ], [3.35750248], [2.65199945], [2.02298999], [1.59444971], [0.80865185], [0.77579791]]) z1 = jnp.array([[ 1.58559484], [ 3.79094097], [-0.52712522], [-1.0178286 ], [-3.51076985], [ 1.30108161], [-1.29824303], [-0.19209007], [ 0.37451138], [-2.33619987]]) init = z[0] distance = 2.6 new = init - distance @jit def test(): # 生成符合条件的布尔掩码 mask = z >= new # 反转掩码,将原数组最后一个符合条件的位置转为反转数组的第一个符合条件位置 flipped_mask = jnp.flip(mask, axis=0) # argmax返回反转数组中第一个True的索引(True对应数值1,优先级高于False的0) flipped_idx = jnp.argmax(flipped_mask) # 计算原数组中的对应索引 original_idx = mask.shape[0] - flipped_idx - 1 # 取出z1中对应的值 return z1[original_idx] # 执行验证 print(test()) # 输出: [-3.51076985]
原理说明
jnp.argmax返回标量索引,形状静态可确定,符合JIT对输入输出形状静态已知的要求;- 通过反转掩码数组,将“找最后一个符合条件的位置”转化为“找第一个符合条件的位置”,避免了
jnp.where返回动态长度数组的问题; - 所有操作均为JAX原生支持的可追踪操作,能被JIT正确编译优化。
内容的提问来源于stack exchange,提问作者C_van3
相关产品推荐
相关产品推荐

