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

如何在不提前导入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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 20:45:06