为元组数组创建布尔掩码:Numpy对象数组全False问题排查
解决Numpy Object数组与元组比较返回全False的问题
这问题我之前也碰到过,核心原因是Numpy对序列类型(比如你的元组template)的比较逻辑和你预期的不一样:
当你执行my_array == template时,Numpy会把这个长度为3的元组当成一个1维数组来处理,然后通过广播机制和你的3x3数组做逐元素比较——也就是说,它会拿my_array[i,j]这个完整元组,去和template[j](元组里的单个元素,比如第0列是字符串"Apple",第1列是"Orange",第2列是5.0)做比较,元组和单个字符串/浮点数当然不相等,结果自然全是False。
而你用元素级比较(比如my_array[0,0] == template)时,是直接比较两个元组对象,所以能得到正确的True。
下面给你几个可行的解决方案:
方案1:用np.vectorize实现元素级比较
vectorize会帮你把单个元素的比较逻辑应用到数组的每个元素上:
import numpy as np template = ('Apple', 'Orange', 5.0) my_array = np.array([None] * 9).reshape((3,3)) for i in range(my_array.shape[0]): for j in range(my_array.shape[1]): my_array[i, j] = template # 定义向量化的比较函数 compare = np.vectorize(lambda x: x == template) mask = compare(my_array) print(mask) # 输出:[[ True True True] # [ True True True] # [ True True True]]
方案2:列表推导式+转Numpy数组
如果你觉得vectorize不够直观,用列表推导式先处理成Python列表,再转回Numpy数组也很简单:
mask = np.array([[elem == template for elem in row] for row in my_array])
这种方式可读性强,对于小尺寸的数组完全够用。
方案3:利用np.apply_along_axis(适合熟悉Numpy轴操作的场景)
mask = np.apply_along_axis(lambda row: [x == template for x in row], axis=1, arr=my_array)
它会沿着行轴(axis=1)遍历每个元素,执行比较逻辑。
额外提示
如果你创建数组时直接用元组初始化,而不是先填None再赋值,代码会更简洁:
my_array = np.array([[template]*3 for _ in range(3)], dtype=object)
这样得到的数组和你之前的完全一样,但避免了嵌套循环赋值。
内容的提问来源于stack exchange,提问作者Oxana Verkholyak
相关产品推荐
相关产品推荐

