You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.02 05:44:51