使用自定义元类时PyTorch Lightning save_hyperparameters报KeyError: 'self'
问题:LightningDataModule自定义元类导致save_hyperparameters抛出KeyError: 'self'
在使用LightningDataModule时,希望无论是否子类化,_after_init方法都能在实例完全初始化后仅执行一次。为此实现了自定义元类_InitMeta,重写__call__方法以在实例创建后调用_after_init,但创建最终子类实例时,save_hyperparameters()内部出现KeyError: 'self'错误。
代码片段
from typing import Any from lightning import LightningDataModule class _InitMeta(type): def __call__( cls: Any, *args: Any, **kwargs: Any ) -> Any: instance = super().__call__(*args, **kwargs) # Create the instance if hasattr(instance, "_after_init"): instance._after_init(**kwargs) # Call the method if defined return instance class A(LightningDataModule, metaclass=_InitMeta): def __init__(self, *args, **kwargs): self.save_hyperparameters() self.a = 1 self.b = 2 super().__init__(*args, **kwargs) def print_ab(self, **kwargs: Any): print("in print ab") if kwargs.get("flag", False): print("flag is set to False") print("some other logic") else: print(self.a, self.b) def _after_init(self, **kwargs): """Called only once after full initialization.""" self.print_ab(**kwargs) class B(A): def __init__(self, **kwargs): super().__init__(**kwargs) self.a += 1 self.b += 2 class C(B): def __init__(self, **kwargs): super().__init__(**kwargs) self.a += 1 self.b += 2 if __name__ == "__main__": print("Creating C instance:") c = C() # Should print 3, 6 only once print("\nCreating B instance:") b = B() # Should print 2, 4 only once print("\nCreating A instance:") a = A() # Should print 1, 2 only once
错误输出
Creating C instance: Traceback (most recent call last): File "G:\github-aditya0by0\python-chebai\test.py", line 48, in <module> c = C() # Should print 3, 6 only once File "G:\github-aditya0by0\python-chebai\test.py", line 10, in __call__ instance = super().__call__(*args, **kwargs) # Create the instance File "G:\github-aditya0by0\python-chebai\test.py", line 41, in __init__ super().__init__(**kwargs) File "G:\github-aditya0by0\python-chebai\test.py", line 34, in __init__ super().__init__(**kwargs) File "G:\github-aditya0by0\python-chebai\test.py", line 18, in __init__ self.save_hyperparameters() File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\core\mixins\hparams_mixin.py", line 112, in save_hyperparameters save_hyperparameters(self, *args, ignore=ignore, frame=frame) File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\utilities\parsing.py", line 165, in save_hyperparameters for local_args in collect_init_args(frame, [], classes=(HyperparametersMixin,)): File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\utilities\parsing.py", line 135, in collect_init_args return collect_init_args(frame.f_back, path_args, inside=True, classes=classes) File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\utilities\parsing.py", line 135, in collect_init_args return collect_init_args(frame.f_back, path_args, inside=True, classes=classes) File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\utilities\parsing.py", line 135, in collect_init_args return collect_init_args(frame.f_back, path_args, inside=True, classes=classes) File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\utilities\parsing.py", line 131, in collect_init_args local_self, local_args = _get_init_args(frame) File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\utilities\parsing.py", line 97, in _get_init_args local_args = {k: local_vars[k] for k in init_parameters} File "G:\anaconda3\envs\env_chebai\lib\site-packages\lightning\pytorch\utilities\parsing.py", line 97, in <dictcomp> local_args = {k: local_vars[k] for k in init_parameters} KeyError: 'self'
环境
- PyTorch Lightning: 2.1.2
- Python: 3.10.14
- Torch: 2.5.1
问题原因
Lightning的save_hyperparameters方法通过解析调用栈帧来获取__init__方法的局部变量,自定义元类的__call__方法改变了默认的实例创建调用栈结构,导致栈帧解析时无法找到self变量,从而抛出KeyError。
解决方案
方案1:调整元类实现,避免干扰栈帧解析
修改元类,通过包装__init__方法而非重写__call__来注入_after_init调用,确保栈帧结构符合Lightning的预期:
from typing import Any import functools from lightning import LightningDataModule class _InitMeta(type): def __new__(cls, name, bases, attrs): # 包装__init__方法 original_init = attrs.get("__init__") if original_init is not None: @functools.wraps(original_init) def wrapped_init(self, *args, **kwargs): original_init(self, *args, **kwargs) # 所有__init__执行完成后调用_after_init,仅一次 if not hasattr(self, "_after_init_called"): if hasattr(self, "_after_init"): self._after_init(**kwargs) setattr(self, "_after_init_called", True) attrs["__init__"] = wrapped_init return super().__new__(cls, name, bases, attrs) class A(LightningDataModule, metaclass=_InitMeta): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) # 先调用父类__init__,符合Lightning初始化规范 self.save_hyperparameters() self.a = 1 self.b = 2 def print_ab(self, **kwargs: Any): print("in print ab") if kwargs.get("flag", False): print("flag is set to False") print("some other logic") else: print(self.a, self.b) def _after_init(self, **kwargs): """Called only once after full initialization.""" self.print_ab(**kwargs) class B(A): def __init__(self, **kwargs): super().__init__(**kwargs) self.a += 1 self.b += 2 class C(B): def __init__(self, **kwargs): super().__init__(**kwargs) self.a += 1 self.b += 2 if __name__ == "__main__": print("Creating C instance:") c = C() # 输出3, 6 print("\nCreating B instance:") b = B() # 输出2, 4 print("\nCreating A instance:") a = A() # 输出1, 2
方案2:直接指定save_hyperparameters的栈帧
在调用save_hyperparameters时,手动指定起始栈帧,跳过元类的__call__方法:
import inspect class A(LightningDataModule, metaclass=_InitMeta): def __init__(self, *args, **kwargs): # 指定当前__init__的栈帧,让Lightning从这里开始解析 self.save_hyperparameters(frame=inspect.currentframe()) self.a = 1 self.b = 2 super().__init__(*args, **kwargs)
这种方法保留元类的__call__实现,但需要修改save_hyperparameters的调用方式,适合不想改动元类逻辑的场景。
内容的提问来源于stack exchange,提问作者Aditya Khedekar
相关产品推荐
相关产品推荐

