在JAX中实现获取数组每行首个未掩码元素索引的最简方案
用JAX实现每行首个未掩码元素的索引获取
原NumPy代码依赖numpy.ma.notmasked_edges获取每行首个未掩码元素的索引,但JAX不支持numpy.ma模块,这里提供基于JAX原生函数的最简实现:
核心实现思路
利用jax.numpy.argmax定位每行第一个未掩码的位置:将原掩码反转(~mask)后,未掩码位置为True,argmax会返回每行第一个True的索引,正好对应目标需求。
完整代码示例
import jax.numpy as jnp from jax import random # 生成测试数据(和原NumPy代码逻辑对齐) key = random.PRNGKey(42) my_array = random.uniform(key, shape=(5,5)) mask = (my_array < 0.5) # 获取每行首个未掩码元素的索引 first_unmasked_indices = jnp.argmax(~mask, axis=1) # 可选:对齐原NumPy.ma行为,全掩码行返回-1 first_unmasked_indices = jnp.where(jnp.any(~mask, axis=1), first_unmasked_indices, -1)
代码说明
~mask:反转掩码数组,将未被掩码的位置标记为Truejnp.argmax(~mask, axis=1):沿行维度(axis=1)查找第一个True的索引,即每行首个未掩码元素的位置- 额外的
jnp.where处理:如果某一行所有元素都被掩码,将索引设为-1,和numpy.ma.notmasked_edges的返回行为完全对齐
内容的提问来源于stack exchange,提问作者Ben
相关产品推荐
相关产品推荐

