请求解析numpy where函数:文档示例的易懂拆解说明
把numpy.where拆得明明白白
我太懂这种对着文档抓耳挠腮的感觉了!numpy.where单看简单比较的用法确实好理解,但文档里的多参数示例总让人摸不着头脑,我来一步步给你拆解清楚。
先回顾你已经懂的:单参数用法
你熟悉的np.where(value > otherValue)属于单参数模式,它的作用很直接:
返回所有满足条件的元素的索引位置,结果是一个元组,每个元素对应数组某一维度的索引数组。
举个简单例子:
import numpy as np arr = np.array([1, 3, 5, 2, 4]) # 找所有大于3的元素的索引 indices = np.where(arr > 3) print(indices) # 输出:(array([2, 4]),)
这个结果表示,数组里第2位(索引从0开始)和第4位的元素满足条件。
重点攻克:三参数的np.where(condition, x, y)
这是文档里最容易让人懵的部分,其实核心规则超简单:
对数组的每一个位置,满足condition就取x对应位置的值,不满足就取y对应位置的值。
这里的x和y可以是单个标量,也可以是和condition同形状的数组,numpy会自动处理广播(不用手动对齐形状)。
示例1:x和y是标量(最常用的替换场景)
比如把数组里大于3的元素换成10,其余换成0:
result = np.where(arr > 3, 10, 0) print(result) # 输出:array([ 0, 0, 10, 0, 10])
逻辑很直观:逐个检查arr的元素,符合条件就用10替换,不符合就用0替换。
示例2:x和y是同形状数组
如果我们想根据条件从两个不同数组里选值,也完全可以:
x = np.array([10, 20, 30, 40, 50]) # 和arr同形状 y = np.array([-1, -2, -3, -4, -5]) result = np.where(arr > 3, x, y) print(result) # 输出:array([-1, -2, 30, -4, 50])
解释:
- arr[0]=1不大于3 → 取y[0]=-1
- arr[2]=5大于3 → 取x[2]=30
- 以此类推,每个位置都按条件二选一。
示例3:多维数组的情况
where对多维数组同样生效,逻辑是逐元素判断:
arr_2d = np.array([[1, 6], [3, 4]]) result = np.where(arr_2d > 3, 'big', 'small') print(result) # 输出: # array([['small', 'big'], # ['small', 'big']], dtype='<U5')
容易混淆的关键点:两种模式的返回值不同
一定要区分开:
- 单参数模式:返回索引元组(用来定位满足条件的元素)
- 三参数模式:返回和输入同形状的新数组(直接生成替换后的结果)
进阶小技巧:多条件嵌套
如果需要多分支判断,可以嵌套where,比如把数组分成三个等级:
result = np.where(arr > 4, 'large', np.where(arr > 2, 'medium', 'small')) print(result) # 输出:array(['small', 'medium', 'large', 'small', 'medium'], dtype='<U6')
逻辑是:先判断是否>4,是就返回'large';否则进入内层where,判断是否>2,是就返回'medium',否则返回'small'。
这样拆解下来是不是清晰多了?如果还有文档里的具体示例搞不懂,随时提出来我再给你拆~
内容的提问来源于stack exchange,提问作者wilson_smyth
相关产品推荐
相关产品推荐

