【发布时间】:2023-03-20 12:30:02
【问题描述】:
如果我创建一个包含 Numpy ndarray 的 Python 数据类,我将无法再使用自动生成的 __eq__。
import numpy as np
@dataclass
class Instr:
foo: np.ndarray
bar: np.ndarray
arr = np.array([1])
arr2 = np.array([1, 2])
print(Instr(arr, arr) == Instr(arr2, arr2))
ValueError:具有多个元素的数组的真值不明确。使用 a.any() 或 a.all()
这是因为ndarray.__eq__ 有时 通过比较a[0] 和b[0],以此类推,直到2 中较长的一个,返回一个ndarray 的真值。这个非常复杂且不直观,实际上只有当数组的形状不同或具有不同的值或其他东西时才会引发错误。
我如何安全地比较持有 Numpy 数组的 @dataclasses?
@dataclass 对__eq__ 的实现是使用eval() 生成的。它的源代码从堆栈跟踪中丢失,无法使用inspect 查看,但它实际上使用了一个元组比较,它调用了bool(foo)。
import dis
dis.dis(Instr.__eq__)
摘录:
3 12 LOAD_FAST 0 (self) 14 LOAD_ATTR 1 (foo) 16 LOAD_FAST 0 (self) 18 LOAD_ATTR 2 (bar) 20 BUILD_TUPLE 2 22 LOAD_FAST 1 (other) 24 LOAD_ATTR 1 (foo) 26 LOAD_FAST 1 (other) 28 LOAD_ATTR 2 (bar) 30 BUILD_TUPLE 2 32 COMPARE_OP 2 (==) 34 RETURN_VALUE
【问题讨论】:
-
您可以在
Instr上编写自己的__eq__方法,您可以覆盖任何自动生成的方法。只需抓住ValueError并实现您自己的附加逻辑。 -
备案,数据类
.__eq__源在这里github.com/python/cpython/blob/3.7/Lib/dataclasses.py#L884
标签: python numpy python-dataclasses