使用numpy自定义数据类型时,数组相等比较不符合预期的问题
问题:Numpy自定义类型数组相等比较未返回预期的Eq对象
问题原因
Numpy对==运算符对应的np.equal通用函数(ufunc)有特殊处理逻辑:它会强制将元素级__eq__方法的返回值转换为布尔类型——这是因为Numpy默认期望比较操作生成布尔数组,用于掩码、条件筛选等常规数值计算场景。而加法+对应的np.add ufunc则不会做强制类型转换,直接将__add__返回的Add对象存入object类型的数组中,因此加法能正常工作。
解决方案
方法一:重载__array_ufunc__拦截相等比较
通过在基类Expr中重载__array_ufunc__方法,拦截np.equal操作,手动生成包含Eq对象的数组:
import numpy as np class Expr: def __add__(self, other): return Add(self, other) def __eq__(self, other): return Eq(self, other) def __array_ufunc__(self, ufunc, method, *inputs, **kwargs): # 拦截相等比较的ufunc调用 if ufunc is np.equal and method == '__call__': left, right = inputs # 创建与输入同形状的object类型数组 result = np.empty_like(left, dtype=object) # 遍历每个元素生成Eq对象 for idx in np.ndindex(left.shape): result[idx] = Eq(left[idx], right[idx]) return result # 其他ufunc操作按默认逻辑处理 return NotImplemented class Variable(Expr): def __init__(self, name): self.name = name def __repr__(self): return self.name class Operator(Expr): def __init__(self, left, right): self.left = left self.right = right def __repr__(self): return f'{self.__class__.__name__}({self.left}, {self.right})' class Add(Operator): ... class Eq(Operator): ... if __name__ == '__main__': arr1 = np.array([Variable('v1'), Variable('v2')]) arr2 = np.array([Variable('v3'), Variable('v4')]) print(arr1 + arr2) print(arr1 == arr2)
运行后输出:
[Add(v1, v3) Add(v2, v4)] [Eq(v1, v3) Eq(v2, v4)]
方法二:使用自定义比较函数替代==运算符
如果不想修改类的底层逻辑,可以直接写一个工具函数来生成Eq对象数组:
def array_eq(arr_a, arr_b): result = np.empty_like(arr_a, dtype=object) for idx in np.ndindex(arr_a.shape): result[idx] = Eq(arr_a[idx], arr_b[idx]) return result # 调用时替换print(arr1 == arr2)为: print(array_eq(arr1, arr2))
这种方式更轻量化,适合不需要全局修改类行为的场景。
内容的提问来源于stack exchange,提问作者Max Berktold
相关产品推荐
相关产品推荐

