Numpy where函数输出含义解析:二维数组条件匹配结果疑问
理解numpy.where()的输出含义
我来给你掰扯清楚这个输出的意思哈!当你只给np.where()传入一个条件(比如x > 5)时,它返回的是所有满足条件的元素的索引坐标集合,而且是按数组的维度分开返回的——二维数组就返回两个数组,分别对应行和列的索引;三维数组会返回三个数组,对应三个维度的索引,以此类推。
咱们结合你的例子拆解:
- 你的数组
x是一个3行3列的二维数组,行索引范围是0、1、2,列索引范围也是0、1、2。 - 满足
x > 5的元素是第2行的所有元素(也就是6.、7.、8.),所以这些元素的行索引全部都是2,这就是第一个数组array([2, 2, 2])的含义;而它们的列索引分别是0、1、2,对应第二个数组array([0, 1, 2])。
把两个数组的元素一一配对,就是每个符合条件的元素的精确坐标:
(2, 0)→x[2, 0] = 6.(2, 1)→x[2, 1] = 7.(2, 2)→x[2, 2] = 8.
再举个小例子帮你巩固:如果条件换成x > 3,np.where(x > 3)会返回(array([1, 1, 2, 2, 2]), array([1, 2, 0, 1, 2])),对应的元素就是4.、5.、6.、7.、8.,你可以自己对应坐标验证下~
内容的提问来源于stack exchange,提问作者user1050619
相关产品推荐
相关产品推荐

