如何将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
相关产品推荐
相关产品推荐

