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

在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:反转掩码数组,将未被掩码的位置标记为True
  • jnp.argmax(~mask, axis=1):沿行维度(axis=1)查找第一个True的索引,即每行首个未掩码元素的位置
  • 额外的jnp.where处理:如果某一行所有元素都被掩码,将索引设为-1,和numpy.ma.notmasked_edges的返回行为完全对齐

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 22:12:39