【问题标题】:How to deserialize a PyTorch saved model with private methods inside a class?如何使用类中的私有方法反序列化 PyTorch 保存的模型?
【发布时间】:2020-04-09 19:03:47
【问题描述】:

我使用 PyTorch 保存方法来序列化一堆基本对象。其中,有一个类在同一个类的 __init__ 中引用了一个私有方法。现在,在序列化之后,我无法反序列化(unpickle)文件,因为在类外部无法访问私有方法。知道如何解决或绕过它吗?我需要恢复保存到该类属性中的数据。

  File ".conda/envs/py37/lib/python3.7/site-packages/IPython/core/interactiveshell.py", line 3331, in run_code
    exec(code_obj, self.user_global_ns, self.user_ns)
  File "<ipython-input-1-a5666d77c70f>", line 1, in <module>
    torch.load("snapshots/model.pth", map_location='cpu')
  File ".conda/envs/py37/lib/python3.7/site-packages/torch/serialization.py", line 529, in load
    return _legacy_load(opened_file, map_location, pickle_module, **pickle_load_args)
  File ".conda/envs/py37/lib/python3.7/site-packages/torch/serialization.py", line 702, in _legacy_load
    result = unpickler.load()
AttributeError: 'Trainer' object has no attribute '__iterator'
  • EDIT-1:

这里有一段代码会产生我现在面临的问题。

import torch

class Test:
    def __init__(self):
        self.a = min
        self.b = max
        self.c = self.__private  # buggy

    def __private(self):
        return None

test = Test()

torch.save({"test": test}, "file.pkl")
torch.load("file.pkl")

但是,如果您从方法中删除私有属性,则不会出现任何错误。

import torch

class Test:
    def __init__(self):
        self.a = min
        self.b = max
        self.c = self.private  # not buggy

    def private(self):
        return None

test = Test()

torch.save({"test": test}, "file.pkl")
torch.load("file.pkl")

【问题讨论】:

    标签: python serialization deployment pytorch pickle


    【解决方案1】:

    这个问题和Python multiprocessing - mapping private method类似,但是因为悬赏而不能被标记为重复。

    该问题源于 Python 错误跟踪器上的这个未解决问题:Objects referencing private-mangled names do not roundtrip properly under pickling,并且与 pickle 处理名称修改的方式有关。有关此答案的更多详细信息:https://stackoverflow.com/a/57497698/6352677

    此时,唯一的解决方法是不使用 __init__ 中的私有方法。

    【讨论】:

      【解决方案2】:

      这个问题是由于 name mangling 造成的——解释器以下面的方式更改变量的名称,这使得以后扩展类时更难产生冲突。在哪里

      self.__private
      

      已更改为 (self._className__privateMethodName)

      self._Test__private
      

      由于 name mangling 不适用于 dunder,其中名称必须以双下划线开头和结尾。

      所以,为了避免 name mangling 在末尾添加两个下划线。

      下面的 sn-p 应该可以解决您的问题。

      import torch
      
      class Test:
          def __init__(self):
              self.a = min
              self.b = max
              self.c = self.__private__
      
          def __private__(self):
              return None
      
      test = Test()
      
      torch.save({"test": test}, "file.pkl")
      torch.load("file.pkl")
      

      【讨论】:

      • OP 的问题是关于__init__ 中的序列化使用私有方法,也就是名称修改。这里__private__ 没有名字修饰,所以不是私有的,尽管有名字。
      • 如果我们只能使用__private 方法,我们就无法通过序列化来实现上述代码。
      • 对。所以你对这个问题的回答是“它不能用私有方法来完成”,这也是我的回答。但是您的答案中没有明确说明,而是说“添加尾随 ”,我觉得这很令人困惑。如果您重命名 __private 变量以使其公开,请明确并称其为公开,而不是误导性的 __private(它是公开的,并且是一种特殊的方法/dunder——出于什么原因?)。
      猜你喜欢
      • 2023-03-06
      • 2014-12-13
      • 2021-06-23
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2017-11-21
      • 2020-01-26
      相关资源
      最近更新 更多