如何解决元类动态生成的Mixin类实例Pickle序列化报错问题
报错原因
pickle序列化自定义类实例时,默认要求类定义存在于所属模块的顶层命名空间中。你通过元类动态创建的混合类是运行时临时生成的,没有注册到模块全局命名空间,序列化时pickle无法通过名称查找到对应类,因此抛出异常。
修复方案
需要在元类创建动态类的逻辑中补充两个处理:
- 把动态生成的混合类注册到所属模块的全局命名空间,供pickle查找
- 为动态类实现
__reduce__方法,明确指定序列化和反序列化的规则
修改后的完整可运行代码如下:
from abc import ABCMeta import sys import pickle class AutoMixinMeta(ABCMeta): """ Helps us conditionally include Mixins, which is useful if we want to switch between different combinations of models (ex. SBERT with Doc Embedding, RoBERTa with positional embeddings). class Sub(metaclass = AutoMixinMeta): def __init__(self, name): self.name = name """ def __call__(cls, *args, **kwargs): mixin = None try: mixin = kwargs.pop('mixin') if isinstance(mixin, list): mixin_names = list(map(lambda x: x.__name__, mixin)) mixin_name = '.'.join(mixin_names) cls_list = tuple(mixin + [cls]) else: mixin_name = mixin.__name__ cls_list = tuple([mixin, cls]) name = "{}With{}".format(cls.__name__, mixin_name) cls = type(name, cls_list, dict(cls.__dict__)) # 新增:将动态类注册到模块全局命名空间,供pickle查找 mod = sys.modules[cls.__module__] if not hasattr(mod, name): setattr(mod, name, cls) # 新增:实现__reduce__方法指定序列化规则 def __reduce__(self): return (Mixer, [], {"mixin": mixin, **self.__dict__}) cls.__reduce__ = __reduce__ except KeyError: pass return type.__call__(cls, *args, **kwargs) class Mixer(metaclass = AutoMixinMeta): """ Class to mix different elements in. a = Mixer(config=config, mixin=[A, B, C]) """ pass class A(): pass class B(): pass config={ 'test_a': True, 'test_b': True } def get_mixins(config): mixins = [] if config['test_a']: mixins.append(A) if config['test_b']: mixins.append(B) return mixins to_mix = get_mixins(config) c = Mixer(mixin=to_mix) # 序列化 pickle.dump(c, open('test.pkl', 'wb')) # 反序列化测试 d = pickle.load(open('test.pkl', 'rb')) print(type(d)) # 输出 <class '__main__.MixerWithA.B'> print(isinstance(d, A)) # 输出 True print(isinstance(d, B)) # 输出 True
内容的提问来源于stack exchange,提问作者Alex Spangher
相关产品推荐
相关产品推荐

