Python __getattr__与torch.nn.Module结合引发无限递归问题排查
PyTorch Module包装类的RecursionError问题分析与修复
问题根源
你的包装类继承了torch.nn.Module,而PyTorch的nn.Module重写了__setattr__和__getattr__方法,导致属性访问逻辑和普通Python对象存在本质差异:
- 属性存储规则:当给
nn.Module子类实例赋值一个nn.Module类型对象(比如示例中的Linear层)时,nn.Module的__setattr__会自动将该对象存入实例的_modules字典,而非普通的__dict__属性字典。 - 无限递归触发:你的
__getattr__方法中尝试访问self.instance,但由于instance不在__dict__中,Python会再次调用你的__getattr__来获取instance属性,形成无限循环。你尝试的self.__dict__['instance']无效,是因为这个键根本不存在于__dict__中。
访问b.test_attribute的递归流程:
- 检查
b的_parameters、_modules、_buffers,找不到test_attribute - 调用自定义
__getattr__('test_attribute') - 方法内尝试访问
self.instance,因instance不在__dict__中,再次触发__getattr__('instance') - 重复上述步骤,直到触发递归深度限制
修复方案
以下两种方案均可解决问题,可根据需求选择:
方案1:绕过nn.Module的__setattr__,直接存入__dict__
在__init__中直接将instance写入实例的__dict__,避免被PyTorch的属性逻辑拦截:
import torch class MyWrapper(torch.nn.Module): def __init__(self, instance): super().__init__() # 直接写入__dict__,绕过nn.Module的__setattr__ self.__dict__['instance'] = instance def __getattr__(self, name): print("trace", name) return getattr(self.instance, name)
方案2:在__getattr__中优先处理instance属性
在自定义__getattr__中,先判断目标属性是否为instance,直接从_modules中取出(PyTorch已将其存在此处),再处理其他属性:
import torch class MyWrapper(torch.nn.Module): def __init__(self, instance): super().__init__() self.instance = instance def __getattr__(self, name): if name == 'instance': # 从_modules中直接获取,避免递归 return self._modules['instance'] print("trace", name) return getattr(self.instance, name)
验证修复效果
运行原测试代码,两种方案均能正常输出:
# 修复后的测试代码 net = torch.nn.Linear(12, 12) net.test_attribute = "hello world" b = MyWrapper(net) print(b.test_attribute) # trace test_attribute\n hello world print(b.instance) # 输出Linear(12, 12)
补充说明
如果你的包装类不需要使用PyTorch模块的特性(比如参数管理、设备迁移等),最简单的方式是让MyWrapper不继承torch.nn.Module,原代码即可直接正常运行。
内容的提问来源于stack exchange,提问作者GooseFromHell
相关产品推荐
相关产品推荐

