【问题标题】:PyTorch Module with attrs cannot get parameter list带有 attrs 的 PyTorch 模块无法获取参数列表
【发布时间】:2019-12-09 00:07:54
【问题描述】:

attr 的包以某种方式破坏了 pytorch 的 parameter() 模块方法。我想知道是否有人有任何变通方法或解决方案,以便这两个包可以无缝集成?

如果没有,关于将问题发布到哪个 github 的任何建议?我的直觉是将其发布到 attr 的 github 上,但堆栈跟踪几乎与 pytorch 的代码库完全相关。

Python 3.7.3
attrs== 19.1.0
torch==1.1.0.post2
torchvision==0.3.0
import attr
import torch


class RegularModule(torch.nn.Module):
    pass

@attr.s
class AttrsModule(torch.nn.Module):
    pass


module = RegularModule()
print(list(module.parameters()))

module = AttrsModule()
print(list(module.parameters()))

实际输出为:

$python attrs_pytorch.py
[]
Traceback (most recent call last):
  File "attrs_pytorch.py", line 18, in <module>
    print(list(module.parameters()))
  File "/usr/local/anaconda3/envs/bgg/lib/python3.7/site-packages/torch/nn/modules/module.py", line 814, in parameters
    for name, param in self.named_parameters(recurse=recurse):
  File "/usr/local/anaconda3/envs/bgg/lib/python3.7/site-packages/torch/nn/modules/module.py", line 840, in named_parameters
    for elem in gen:
  File "/usr/local/anaconda3/envs/bgg/lib/python3.7/site-packages/torch/nn/modules/module.py", line 784, in _named_members
    for module_prefix, module in modules:
  File "/usr/local/anaconda3/envs/bgg/lib/python3.7/site-packages/torch/nn/modules/module.py", line 975, in named_modules
    if self not in memo:
TypeError: unhashable type: 'AttrsModule'

预期的输出是:

$python attrs_pytorch.py
[]
[]

【问题讨论】:

  • 为什么要做空班?
  • 好问题。空类只是一个最小的例子。它们表明不是任何其他类方法造成问题,而是包本身。
  • 我想肯定有特殊的 init 以特殊的方式调用 super().init,可能是因为这个。是的,这是一个设计失误。

标签: python python-3.x pytorch python-attrs


【解决方案1】:

您可以使用一种解决方法并使用dataclasses(您应该使用它,因为它在标准Python库中,因为您显然正在使用3.7)。虽然我认为简单的__init__ 更具可读性。可以使用attrs 库(禁用散列)来做类似的事情,如果可能的话,我只是更喜欢使用标准库的解决方案。

原因(如果您设法处理与散列相关的错误)是您正在调用 torch.nn.Module.__init__(),它会生成 _parameters 属性和其他特定于框架的数据。

首先用dataclasses解决散列问题:

@dataclasses.dataclass(eq=False)
class AttrsModule(torch.nn.Module):
    pass

这解决了hashing 的问题,正如documentation 所述,关于hasheq 的部分:

默认情况下,dataclass() 不会隐式添加 hash() 方法 除非这样做是安全的。

这是 PyTorch 需要的,因此该模型可以在 C++ 支持中使用(如果我错了,请纠正我),此外:

如果 eq 为假,hash() 将保持不变,这意味着 将使用超类的 hash() 方法(如果超类是对象,这意味着它将回退到基于 id 的散列)。

所以您可以使用 torch.nn.Module __hash__ 函数(如果出现任何进一步的错误,请参阅数据类文档)。

这会给你留下错误:

AttributeError: 'AttrsModule' object has no attribute '_parameters'

因为torch.nn.Module 构造函数没有被调用。快速而肮脏的修复:

@dataclasses.dataclass(eq=False)
class AttrsModule(torch.nn.Module):
    def __post_init__(self):
        super().__init__()

__post_init__ 是在__init__ 之后调用的函数(谁会猜到),您可以在其中初始化特定于 Torch 的参数。

不过,我建议反对同时使用这两个模块。例如,您正在使用您的代码破坏 PyTorch 的 __repr__,因此应将 repr=False 传递给 dataclasses.dataclass 构造函数,这会给出最终代码(我希望消除库之间的明显冲突):

import dataclasses

import torch


class RegularModule(torch.nn.Module):
    pass


@dataclasses.dataclass(eq=False, repr=False)
class AttrsModule(torch.nn.Module):
    def __post_init__(self):
        super().__init__()


module = RegularModule()
print(list(module.parameters()))

module = AttrsModule()
print(list(module.parameters()))

有关attrs 的更多信息,请参阅hynek 答案和他的博文。

【讨论】:

  • 哇,非常感谢您如此周到周到的回复!我不知道dataclasses,一定会调查的!
  • 是什么让你说“这个 stdlib 和 attrs 一样,甚至更多。”?数据类是属性的严格子集。
  • @hynek 如果是这样,您可以编辑我的答案或添加评论以扩展它吗?我对attrs 不太熟悉,所以您的意见非常有用,谢谢。我已为您的答案添加了链接,如有必要,将进行扩展。
  • 编辑了我的答案以纳入您的评论,感谢分享,不想传播 FUD,soz。
【解决方案2】:

attrs 有一个关于哈希性的章节,也解释了 Python 中哈希的陷阱:https://www.attrs.org/en/stable/hashing.html

您必须决定哪种行为适合您的具体问题。如需更多一般信息,请查看https://hynek.me/articles/hashes-and-equality/ — 事实证明,散列在 Python 中非常棘手。

【讨论】:

    猜你喜欢
    • 2020-10-04
    • 2022-12-14
    • 1970-01-01
    • 1970-01-01
    • 2019-07-11
    • 2023-03-28
    • 1970-01-01
    • 1970-01-01
    • 2017-10-23
    相关资源
    最近更新 更多