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

Numpy:如何用where()获取元组数组中指定元组的索引

关于numpy.where获取元组数组索引的疑问解答
  • 你对numpy.all(axis=1)的理解完全正确:
    执行a == (1, 2)得到的是2维布尔数组,numpy.all(axis=1)会对每行元素做逻辑与运算,最终生成形状为(3,)的一维布尔数组,对应结果是array([False, True, False]),这一步的行为符合预期。

  • 为什么numpy.where的返回需要用b[0][0]才能拿到目标索引?
    因为numpy.where的返回值是一个元组,元组中的每个元素对应输入数组某一维度的匹配索引集合。针对一维的布尔数组输入,返回的元组里仅包含一个元素——存储匹配索引的一维数组,所以你得到的b实际结构是(array([1]),),而非单独的数组。
    因此必须先通过b[0]取出这个一维索引数组,再用[0]提取其中的单个索引值。

  • 简化索引获取的写法:
    如果你能确定目标元组在数组中仅有一个匹配项,可以用更简洁的方式获取索引:

    import numpy as np
    
    a = np.array([(0, 1), (1, 2), (2, 3)])
    # 方式1:使用argwhere
    idx = np.argwhere(np.all(a == (1, 2), axis=1))[0][0]
    # 方式2:使用flatnonzero
    idx = np.flatnonzero(np.all(a == (1, 2), axis=1))[0]
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 22:22:09