如何在不提前导入torch的前提下定义继承自torch.nn.Module的类?
可行实现方案
以下两种方案都可以实现只在用到MyClass时才加载torch依赖,不需要使用该类的进程导入modules.py时不会触发torch导入:
方案1:使用PEP 562 模块级__getattr__延迟加载(Python 3.7+ 推荐)
该方案符合Python标准规范,调用方可以和普通类一样正常导入使用,无感知:
# modules.py 文件内容 from typing import Any # 仅静态类型检查阶段导入torch,运行时不生效 TYPE_CHECKING = False if TYPE_CHECKING: import torch def __getattr__(name: str) -> Any: if name == "MyClass": # 只有访问MyClass时才导入torch依赖 import torch class MyClass(torch.nn.Module): def __init__(self): super().__init__() # 你的类逻辑实现 ... # 缓存类到模块全局,下次访问无需重复定义 globals()["MyClass"] = MyClass return MyClass raise AttributeError(f"module {__name__} has no attribute {name}")
使用方式和普通类完全一致:
# 用到MyClass的进程中直接导入即可 from modules import MyClass obj = MyClass()
方案2:类工厂函数(全Python版本兼容)
如果需要兼容Python 3.7以下版本,可以用工厂函数封装类定义:
# modules.py 文件内容 TYPE_CHECKING = False if TYPE_CHECKING: import torch from typing import Type def get_my_class() -> "Type[torch.nn.Module]": # 调用函数时才导入torch import torch class MyClass(torch.nn.Module): def __init__(self): super().__init__() # 你的类逻辑实现 ... return MyClass
使用方式:
# 需要用到时调用工厂函数获取类再实例化 from modules import get_my_class MyClass = get_my_class() obj = MyClass()
注意事项
- 两种方案都不会影响
modules.py中其他普通类的定义和使用 TYPE_CHECKING常量仅在IDE静态检查、mypy类型校验阶段为真,运行时不会触发torch导入,不影响内存占用
内容的提问来源于stack exchange,提问作者Robin Lobel
相关产品推荐
相关产品推荐

