使用np.where时如何保留NumPy数组的原始结构?
问题描述
示例代码
import numpy as np x = np.transpose(np.array([np.arange(10), np.zeros(10, dtype=int)])) x = np.array([x, x, x]) print("orig: \n", x) print("") print("indexed: \n", x[np.where(np.logical_and(x[..., 0] > 3, x[..., 0] < 7))])
实际输出
orig: [[[0 0] [1 0] [2 0] [3 0] [4 0] [5 0] [6 0] [7 0] [8 0] [9 0]] [[0 0] [1 0] [2 0] [3 0] [4 0] [5 0] [6 0] [7 0] [8 0] [9 0]] [[0 0] [1 0] [2 0] [3 0] [4 0] [5 0] [6 0] [7 0] [8 0] [9 0]]] indexed: [[4 0] [5 0] [6 0] [4 0] [5 0] [6 0] [4 0] [5 0] [6 0]]
期望输出
[[[4 0] [5 0] [6 0]] [[4 0] [5 0] [6 0]] [[4 0] [5 0] [6 0]]]
疑问
- 猜测这是因为
np.where将匹配的最后维度数组放入新数组导致的,是否正确? - 能否通过
np.where实现期望的结果?若可以,如何操作?若不行,有没有更合适的方法?优先使用NumPy向量化函数,而非循环。
解答
关于你的猜测
这个猜测不太准确。实际原因是:np.where返回的是各维度上匹配元素的索引元组,用这些索引取数时,所有符合条件的元素会被提取出来并扁平化(最终变成二维数组,因为每个元素是长度为2的子数组),并非单纯因为“最后维度数组被放入新数组”。
能不能用np.where实现期望结果?
直接用np.where做不到——它的索引逻辑就是提取所有符合条件的元素,没法保留原数组的三维结构。
推荐的向量化解决方案
要保留原数组的前两个维度结构,只筛选第三维度中符合条件的元素,用布尔掩码+维度调整或者np.compress就能搞定,都是纯向量化操作:
方法一:布尔掩码+reshape
import numpy as np x = np.transpose(np.array([np.arange(10), np.zeros(10, dtype=int)])) x = np.array([x, x, x]) # 生成掩码:针对每个子数组的第一个元素判断条件 mask = np.logical_and(x[..., 0] > 3, x[..., 0] < 7) # 提取符合条件的元素后,重塑成期望的三维结构 result = x[mask].reshape(x.shape[0], -1, x.shape[-1]) print("result: \n", result)
x[mask]会取出所有符合条件的元素,得到(9,2)的数组,再通过reshape恢复成(3,3,2)的三维结构。
方法二:用np.compress(更直接)
import numpy as np x = np.transpose(np.array([np.arange(10), np.zeros(10, dtype=int)])) x = np.array([x, x, x]) mask = np.logical_and(x[..., 0] > 3, x[..., 0] < 7) # 沿着第2个维度(axis=1)筛选符合条件的元素 result = np.compress(mask, x, axis=1) print("result: \n", result)
np.compress可以指定维度筛选,直接输出(3,3,2)的结果,不用额外调整形状,更贴合需求。
最终结果
两种方法都会输出你想要的结构:
[[[4 0] [5 0] [6 0]] [[4 0] [5 0] [6 0]] [[4 0] [5 0] [6 0]]]
内容的提问来源于stack exchange,提问作者MattHusz
相关产品推荐
相关产品推荐

