pickle加载属性查找机制及跨模块类反序列化问题解决
关于Pickle跨模块反序列化类的问题
问题场景
我搞不懂pickle加载文件时是怎么查找类属性的,想让不同模块用各自的类定义来完成反序列化,但尝试后失败了。具体场景如下:
原模块代码
mod_a.py
# mod_a.py import pickle class A(object): def __init__(self): self.mod_a = True #------------------------------------------------------------------ def load(file): print('Unpickle',repr(A)) return pickle.load(file) if __name__ == '__main__': # 创建pickle文件 a = A() with open('a.ar','wb') as f: pickle.dump(a,f,protocol=2)
mod_b.py
# mod_b.py import pickle class A(object): def __init__(self): self.mod_b = True #------------------------------------------------------------------ def load(file): print('Unpickle',repr(A)) return pickle.load(file) if __name__ == '__main__': # 创建pickle文件 a = A() with open('b.ar','wb') as f: pickle.dump(a,f,protocol=2)
mod_c.py
import mod_a import mod_b # from mod_a import A # from mod_b import A with open('a.ar','rb') as f: print(mod_a.load(f))
运行报错
先通过运行mod_a.py和mod_b.py生成a.ar和b.ar文件,再运行mod_c.py时出现如下错误:
Unpickle <class 'mod_a.A'> Traceback (most recent call last): File "C:\proj_py\GTC3\tmp\mod_c.py", line 8, in <module> print(mod_a.load(f)) ^^^^^^^^^^^^^ File "C:\proj_py\GTC3\tmp\mod_a.py", line 10, in load return pickle.load(file) ^^^^^^^^^^^^^^^^^ AttributeError: Can't get attribute 'A' on <module '__main__' from 'C:\proj_py\GTC3\tmp\mod_c.py'>
我原本以为pickle会在load()函数所在的模块(mod_a)里找到A类的定义(print语句也输出了正确的类对象),但实际并没有生效。取消mod_c.py中注释的导入语句能临时解决错误,但会导致用错模块的类,不符合调用mod_a.load(f)就应该用mod_a中A类的预期。
请问能不能修改mod_a和mod_b,让它们的load()函数反序列化时使用对应模块中的类?
解决方案:自定义Unpickler指定类查找逻辑
通过继承pickle.Unpickler并重写find_class方法,强制优先从当前模块查找类定义,即可实现调用mod_a.load()就反序列化为mod_a的A类,调用mod_b.load()就反序列化为mod_b的A类。
修改后的代码
mod_a.py
# mod_a.py import pickle import sys class _Unpickler(pickle.Unpickler): def find_class(self, module, name): if hasattr(sys.modules[__name__],name): return getattr(sys.modules[__name__],name) else: return super(_Unpickler,self).find_class(module, name) class A(object): def __init__(self): self.mod_a = True def load(file): return _Unpickler(file).load() if __name__ == '__main__': a = A() with open('a.ar','wb') as f: pickle.dump(a,f,protocol=2)
mod_b.py
# mod_b.py import pickle import sys class _Unpickler(pickle.Unpickler): def find_class(self, module, name): if hasattr(sys.modules[__name__],name): return getattr(sys.modules[__name__],name) else: return super(_Unpickler,self).find_class(module, name) class A(object): def __init__(self): self.mod_b = True def load(file): return _Unpickler(file).load() if __name__ == '__main__': a = A() with open('b.ar','wb') as f: pickle.dump(a,f,protocol=2)
mod_c.py
# mod_c.py import mod_a import mod_b with open('a.ar','rb') as f: print( mod_a.load(f) ) # print( mod_b.load(f) )
注意事项
此方法并非强制使用A类的初始定义(pickle文件本身并未记录足够的类定义细节),仅明确指定反序列化时使用调用方模块中的A类定义。
内容的提问来源于stack exchange,提问作者Blair
相关产品推荐
相关产品推荐

