使用Enum定义NumPy数组触发广播错误,求Python模块静态数组方案
解决NumPy数组作为Enum值的报错问题及静态不可变数组方案
一、直接将NumPy数组作为Enum值的方法
问题根源是Enum初始化时会用==检查成员值是否重复,但NumPy数组的==是元素级操作,还会尝试广播,形状不同就触发报错。只需自定义Enum类,修改比较逻辑即可解决:
import enum import numpy as np class NumpyEnum(enum.Enum): # 重写相等判断逻辑:成员间用身份比较,避免触发NumPy的广播操作 def __eq__(self, other): if isinstance(other, type(self)): return self is other # 和外部数组比较时,用NumPy的数组相等判断 return np.array_equal(self.value, other) class Foo(NumpyEnum): BAR = np.array([1, 2, 3]) BAZ = np.array([4, 5])
使用示例:
print(Foo.BAR.value) # 输出: [1 2 3] print(Foo.BAR == np.array([1,2,3])) # 输出: True print(Foo.BAR == Foo.BAZ) # 输出: False
二、更优雅的静态不可变NumPy数组声明方案
如果只是需要在模块中存静态不可变数组,没必要用Enum,直接设置数组的writeable标志为False即可:
1. 模块级别直接声明
import numpy as np BAR = np.array([1, 2, 3]) BAR.flags.writeable = False # 设置为不可写 BAZ = np.array([4, 5]) BAZ.flags.writeable = False
尝试修改数组元素会触发报错:ValueError: assignment destination is read-only,保证数组不可变。
2. 类封装(归类管理更清晰)
import numpy as np class StaticArrays: # 初始化时创建并设置不可变数组 _BAR = np.array([1, 2, 3]) _BAR.flags.writeable = False _BAZ = np.array([4, 5]) _BAZ.flags.writeable = False # 用类属性暴露,避免被重新赋值 @classmethod @property def BAR(cls): return cls._BAR @classmethod @property def BAZ(cls): return cls._BAZ
使用时直接通过类访问:
print(StaticArrays.BAR) # 输出: [1 2 3]
内容的提问来源于stack exchange,提问作者Spuu
相关产品推荐
相关产品推荐

