如何实现numpy array多层继承?解决子类属性丢失问题
解决numpy.ndarray多层子类化时自定义属性丢失的问题
问题根源:使用view()转换数组类型时,numpy只会处理数组本身的结构,不会自动复制用户添加的实例属性(比如你的_points)。因此OrderedDataset实例无法继承原Dataset对象的_points属性,导致调用average()时出错。
解决方案:重写__array_finalize__方法传递属性
numpy子类化中,__array_finalize__方法用于在创建新实例(包括视图转换、切片操作等场景)时,从原对象复制属性。修改Dataset类添加该方法,即可自动传递自定义属性:
import numpy as np class Dataset(np.ndarray): def __new__(cls, points: int): dataset = np.zeros(points).view(Dataset) dataset._points = points return dataset def __array_finalize__(self, obj): # 从原对象复制_points属性到新实例 if obj is not None: self._points = getattr(obj, '_points', None) def average(self) -> float: return sum(self) / self._points class OrderedDataset(Dataset): def __new__(cls, points: int): return Dataset(points).view(OrderedDataset) def check_order(self): for i in range(1, len(self)): assert self[i] >= self[i - 1] def test(): d = Dataset(100) assert np.allclose(d.average(), 0) od = OrderedDataset(100) od.check_order() assert np.allclose(od.average(), 0) # 现在正常运行 test()
原理说明
当执行Dataset(points).view(OrderedDataset)时,numpy会创建OrderedDataset实例,并调用其继承自Dataset的__array_finalize__方法,将原Dataset对象的_points属性复制到新的OrderedDataset实例中,确保属性不丢失。
这种方式也适用于后续更多层级的子类化,只要父类正确实现__array_finalize__,子类就能自动继承所需属性。
内容的提问来源于stack exchange,提问作者Miicck
相关产品推荐
相关产品推荐

