如何简化含递归类型与Protocol的Python静态类型标注代码?
简化静态类型标注的方案
可以通过泛型(Generic)来简化这四个Protocol的定义,避免重复代码。利用typing.Generic定义一个带输入、输出类型参数的通用Protocol,再通过不同的类型参数组合出你需要的四种函数类型,最后统一到InitFn的Union中。
修改后的代码如下:
import numpy as np from typing import Any, Iterable, Mapping, Sequence, Union, Protocol, Generic, TypeVar Shape = Sequence[Union[int, np.int32, np.int64]] ShapeTree = Union[Shape, Iterable['ShapeTree'], Mapping[Any, 'ShapeTree']] # 定义两个受约束的类型变量,限定为Shape或ShapeTree的子类型 InputT = TypeVar('InputT', bound=Union[Shape, ShapeTree]) OutputT = TypeVar('OutputT', bound=Union[Shape, ShapeTree]) class InitFnProtocol(Protocol, Generic[InputT, OutputT]): def __call__(self, input_shape: InputT) -> OutputT: ... # 组合出四种所需的函数类型 InitFn = Union[ InitFnProtocol[Shape, Shape], InitFnProtocol[Shape, ShapeTree], InitFnProtocol[ShapeTree, Shape], InitFnProtocol[ShapeTree, ShapeTree] ]
说明
- 用
TypeVar定义的InputT和OutputT限定了输入输出的类型范围,确保只能是Shape或ShapeTree,完全匹配原代码的类型约束。 - 泛型Protocol
InitFnProtocol作为统一模板,通过不同的类型参数组合,生成原本四个独立Protocol对应的函数类型。 - 这种写法大幅减少了重复的类定义,后续如果需要新增输入输出组合,只需在
InitFn的Union中添加新的泛型实例即可,维护性更强。
内容的提问来源于stack exchange,提问作者Miguel Monteiro
相关产品推荐
相关产品推荐

