如何成功pickle.dump含绑定方法的B类实例(关联带paddle.Tensor的A类)
解决pickle序列化包含paddle.Tensor的绑定方法实例问题
问题根源
a.func作为绑定方法,内部持有对实例a的引用,pickle序列化b时会递归序列化a,而a.x是paddle.Tensor对象,无法被pickle直接序列化,因此抛出TypeError: can't pickle tensor object。
方案1:注册paddle.Tensor的全局序列化规则
无需修改原有类结构,通过copyreg为paddle.Tensor注册自定义序列化逻辑,让pickle自动处理Tensor:
import paddle import typing import pickle import copyreg # 注册paddle.Tensor的pickle规则 def pickle_paddle_tensor(tensor): # 序列化时返回构造函数和所需参数(numpy数组) return paddle.to_tensor, (tensor.numpy(),) copyreg.pickle(paddle.Tensor, pickle_paddle_tensor) class A(object): def __init__(self): self.x = paddle.to_tensor([1, 2, 3]) def func(self): pass class B(object): def __init__(self, f: typing.Callable): self.func = f def main(): a = A() b = B(f=a.func) # 成功执行dump with open("./dumpfile", "wb") as f: pickle.dump(b, f) # 验证反序列化 with open("./dumpfile", "rb") as f: loaded_b = pickle.load(f) loaded_b.func() print("加载后实例的Tensor:", loaded_b.func.__self__.x) if __name__ == "__main__": main()
方案2:为类A自定义序列化逻辑
如果不想全局修改Tensor的序列化规则,可在类A中实现__getstate__和__setstate__方法,手动处理Tensor的序列化与恢复:
import paddle import typing import pickle class A(object): def __init__(self): self.x = paddle.to_tensor([1, 2, 3]) def func(self): pass def __getstate__(self): # 序列化时将Tensor转换为numpy数组 state = self.__dict__.copy() state["x"] = self.x.numpy() return state def __setstate__(self, state): # 反序列化时将numpy数组恢复为Tensor self.__dict__.update(state) self.x = paddle.to_tensor(state["x"]) class B(object): def __init__(self, f: typing.Callable): self.func = f def main(): a = A() b = B(f=a.func) with open("./dumpfile", "wb") as f: pickle.dump(b, f) with open("./dumpfile", "rb") as f: loaded_b = pickle.load(f) loaded_b.func() print("加载后实例的Tensor:", loaded_b.func.__self__.x) if __name__ == "__main__": main()
内容的提问来源于stack exchange,提问作者Huo Ke
相关产品推荐
相关产品推荐

