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

如何将nn.Module子类与Protocol类型相交并定义为TypeVar?

实现方案

1. 定义描述构造方法签名的Protocol

先通过Protocol明确类必须具备的构造方法规则——接收dict类型的cfg参数:

from typing import Protocol, Dict
import torch.nn as nn

class ConfigurableModuleProtocol(Protocol):
    def __init__(self, cfg: Dict) -> None:
        ...

2. 定义绑定双特性交集的TypeVar

用TypeVar的bound参数,将类型限制为同时继承nn.Module且符合上述Protocol的类类型:

from typing import TypeVar

ConfigurableModule = TypeVar(
    "ConfigurableModule",
    bound=type[nn.Module] & type[ConfigurableModuleProtocol]
)

这里的type[]表示类本身的类型(而非实例),&用来表达类型交集,确保只有同时满足两个条件的类才能匹配这个TypeVar。

3. 作为Generic类的模板参数使用

将这个TypeVar传入Generic类后,类型检查工具(mypy、PyCharm)会自动校验传入的类是否合法:

from typing import Generic

class ModuleManager(Generic[ConfigurableModule]):
    def __init__(self, module_cls: ConfigurableModule, cfg: Dict):
        # 此处会自动校验module_cls的__init__是否符合签名要求
        self.module = module_cls(cfg)

# 测试:合法类(同时满足两个条件)
class ValidModule(nn.Module):
    def __init__(self, cfg: Dict) -> None:
        super().__init__()
        self.cfg = cfg

# 测试:不合法类(构造方法签名不符合)
class InvalidModule(nn.Module):
    def __init__(self, name: str) -> None:
        super().__init__()
        self.name = name

# 合法调用,无类型错误
manager = ModuleManager(ValidModule, {"lr": 0.01})

# 非法调用,类型检查工具会提示错误
# manager = ModuleManager(InvalidModule, {"lr": 0.01})

关键细节说明

  • 若使用Python 3.9及以下版本,type[X]语法需替换为Type[X](需从typing导入Type),类型交集可以用Union[type[nn.Module], type[ConfigurableModuleProtocol]]替代,但&的语义更精准。
  • Protocol采用结构子类型匹配,无需类显式继承该Protocol,只要构造方法签名符合要求即可被识别。
  • 绑定交集后的TypeVar,能让Generic类在接收类参数时,自动完成双重校验:是否是nn.Module子类、是否有符合要求的构造方法。

内容的提问来源于stack exchange,提问作者Albert

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 01:13:17