如何用numpy保存加载派生类实例且不丢失类引用?
Numpy保存加载派生自ndarray的类实例时丢失类型与属性的解决方法
通过numpy的save/load例程可以保存并加载普通类实例,但在处理派生自np.ndarray的类时,加载后的对象会丢失子类类型,退化为父类numpy.ndarray,同时子类自定义的属性也会丢失,导致访问时报错。
最小可复现示例(MWE)
import numpy as np class SampleClass(object): def __init__(self, a): self.a = a class SampleDerivedClass(np.ndarray): def __new__(cls, a, parameter): obj = a.view(cls) obj.parameter = parameter return obj a=np.array([1,2,3]) x=SampleClass(a) np.save("x",x) xx=np.load("x.npy", allow_pickle=True) print("Saving a class and loading:", type(xx.item())) y=SampleDerivedClass(a,42) y+=a print("Derived class:", type(y),y,y.parameter) np.save("y",y) yy=np.load("y.npy", allow_pickle=True) print("After saving and loading:",type(yy),yy) # Access to `yy.parameter` would result in an error.
实际输出
Saving a class and loading: <class '__main__.SampleClass'> Derived class: <class '__main__.SampleDerivedClass'> [2 4 6] 42 After saving and loading: <class 'numpy.ndarray'> [2 4 6]
此时访问yy.parameter会触发AttributeError。
解决方案
方法1:自定义序列化方法
numpy默认不会保存ndarray派生类的额外属性和类型信息,需要实现__reduce__和__setstate__方法来自定义序列化逻辑:
import numpy as np class SampleDerivedClass(np.ndarray): def __new__(cls, a, parameter): obj = a.view(cls) obj.parameter = parameter return obj def __reduce__(self): # 返回构造函数、构造参数、实例状态的元组 return (self.__class__, (np.array(self),), {'parameter': self.parameter}) def __setstate__(self, state): # 恢复自定义属性 self.parameter = state['parameter']
修改后测试,加载后的对象将保留子类类型和parameter属性:
a=np.array([1,2,3]) y=SampleDerivedClass(a,42) y+=a np.save("y",y) yy=np.load("y.npy", allow_pickle=True) print("After saving and loading:",type(yy),yy,yy.parameter) # 输出:After saving and loading: <class '__main__.SampleDerivedClass'> [2 4 6] 42
方法2:直接使用pickle序列化
numpy的save/load对ndarray派生类的支持有限,直接用pickle可以完整保存对象的类型和属性:
import pickle # 保存对象 with open("y.pkl", "wb") as f: pickle.dump(y, f) # 加载对象 with open("y.pkl", "rb") as f: yy = pickle.load(f) print("After saving and loading:",type(yy),yy,yy.parameter)
方法3:采用组合而非继承
如果业务场景允许,避免直接继承np.ndarray,改用组合方式将数组作为类属性:
class SampleWrapperClass: def __init__(self, arr, parameter): self.arr = arr self.parameter = parameter # 测试示例 a=np.array([1,2,3]) y=SampleWrapperClass(a,42) y.arr +=a np.save("y",y) yy=np.load("y.npy", allow_pickle=True).item() print("After saving and loading:",type(yy),yy.arr,yy.parameter)
内容的提问来源于stack exchange,提问作者Joce
相关产品推荐
相关产品推荐

