Python向量类一维输入函数的类型标注方案求助
一维向量类输入的Python类型标注解决方案
下面提供几种适配需求的方案,你可以根据场景选择:
方案1:组合Sized与Iterable抽象基类
既然需要输入同时支持遍历和获取长度,直接用这两个抽象基类的交集类型即可,覆盖所有符合条件的类型:
from typing import Iterable, Sized def fun(arg1: Iterable[Sized], arg2: Iterable[Sized]) -> None: assert len(arg1) > 0, "length of the arg1 has to be at least 1" assert len(arg1) == len(arg2), "arg1 and arg2 must have the same length" # 后续遍历逻辑
优点是范围宽泛,自动适配所有满足条件的类型;缺点是不够精准,可能会匹配到一些你不需要的类型。
方案2:自定义明确的联合类型别名
直接把允许的类型列出来,用Union组合成别名,类型检查工具(如mypy)能精准识别:
from typing import Union, List, Tuple import numpy as np VectorLike = Union[List, Tuple, np.ndarray] def fun(arg1: VectorLike, arg2: VectorLike) -> None: assert isinstance(arg1, (list, tuple, np.ndarray)), "arg1 must be list, tuple or numpy array" assert isinstance(arg2, (list, tuple, np.ndarray)), "arg2 must be list, tuple or numpy array" assert len(arg1) > 0, "length of the arg1 has to be at least 1" assert len(arg1) == len(arg2), "arg1 and arg2 must have the same length" # 后续遍历逻辑
这个方案最直观,完全符合你对输入类型的限定,适合需要严格控制输入范围的场景。
方案3:扩展Sequence类型兼容numpy数组
虽然np.ndarray原生不属于Sequence,但可以通过注册虚拟子类的方式,让它被isinstance和类型检查工具识别为Sequence:
from collections.abc import Sequence import numpy as np # 全局注册np.ndarray为Sequence的虚拟子类 Sequence.register(np.ndarray) def fun(arg1: Sequence, arg2: Sequence) -> None: assert len(arg1) > 0, "length of the arg1 has to be at least 1" assert len(arg1) == len(arg2), "arg1 and arg2 must have the same length" # 后续遍历逻辑
注意:这个注册是全局生效的,会影响整个程序中isinstance(obj, Sequence)的判断,如果你只在特定模块使用,需要谨慎考虑。
额外优化:一维检查
如果需要确保输入是一维向量,可以在运行时判断中增加维度校验:
# 基于方案2的扩展 from typing import Union, List, Tuple import numpy as np VectorLike = Union[List, Tuple, np.ndarray] def fun(arg1: VectorLike, arg2: VectorLike) -> None: assert isinstance(arg1, (list, tuple, np.ndarray)), "arg1 must be list, tuple or numpy array" assert isinstance(arg2, (list, tuple, np.ndarray)), "arg2 must be list, tuple or numpy array" # 校验numpy数组为一维 if isinstance(arg1, np.ndarray): assert arg1.ndim == 1, "arg1 must be a 1D numpy array" if isinstance(arg2, np.ndarray): assert arg2.ndim == 1, "arg2 must be a 1D numpy array" assert len(arg1) > 0, "length of the arg1 has to be at least 1" assert len(arg1) == len(arg2), "arg1 and arg2 must have the same length"
内容的提问来源于stack exchange,提问作者polkas
相关产品推荐
相关产品推荐

