You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用自定义元类时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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.13 13:14:58