NumPy中where函数用法解析及给定代码输出说明
解释NumPy中np.where()的输出与用法
先看给定的代码及输出:
import numpy as np a = np.array([[1,2],[3,4]]) np.where(a<4)
输出为:(array([0, 0, 1]), array([0, 1, 0]))
代码输出解释
np.where传入布尔条件时,返回的是满足条件的所有元素的坐标索引,按数组维度拆分:
- 第一个数组是行索引,第二个是列索引(因为a是2维数组)
- 逐个核对数组a的元素:
- 1(行0,列0):1<4,满足,对应索引(0,0)
- 2(行0,列1):2<4,满足,对应索引(0,1)
- 3(行1,列0):3<4,满足,对应索引(1,0)
- 4(行1,列1):不满足条件,被忽略
所以最终返回行索引数组[0,0,1],列索引数组[0,1,0],两者组合就是所有符合条件的元素位置。
np.where的三种常用场景
1. 仅传入布尔条件:获取满足条件的元素索引
这就是示例中的用法,返回各维度索引组成的元组,每个维度对应一个索引数组,适用于定位数组中符合要求的元素位置。
2. 条件+替换值:按规则修改数组
格式:np.where(条件, 满足条件时的替换值, 不满足时的替换值)
比如把示例中的数组a里小于4的元素换成-1,其余换成99:
np.where(a < 4, -1, 99) # 输出:array([[-1, -1], # [-1, 99]])
这种用法可以快速完成数组的条件式批量替换。
3. 处理多维数组
np.where支持任意维度的数组,比如3维数组:
b = np.array([[[1,5],[2,6]],[[3,7],[4,8]]]) np.where(b < 5) # 输出:(array([0, 0, 1, 1]), array([0, 1, 0, 1]), array([0, 0, 0, 0]))
返回的三个数组分别对应3维数组的深度、行、列索引,组合起来就是所有满足b<5的元素坐标。
内容的提问来源于stack exchange,提问作者Prasad Nalawade
相关产品推荐
相关产品推荐

