为支持numpy数组的函数加类型提示:解决mypy类型不兼容问题
问题:Mypy拒绝将numpy数组视为
Sequence[float]参数 代码示例与错误信息
import numpy as np from typing import Sequence def compute(x: Sequence[float]) -> bool: # 不修改x的计算逻辑 ... compute(np.linspace(0, 1, 10))
Mypy报错:
Argument 1 to "compute" has incompatible type "ndarray[Any, dtype[floating[Any]]]"; expected "Sequence[float]" [arg-type]
用户疑问
我原以为numpy数组满足Sequence的核心要求(可迭代、可反转、可索引),但Mypy却报错了。这是不是因为numpy数组是可变的,而Sequence更偏向不可变对象?我试过把类型改成Iterable能解决问题,但函数里需要对x做索引操作,所以得找个既支持迭代又支持索引的类型提示方案。
解答
为什么numpy数组不被视为Sequence?
typing.Sequence是Python的抽象基类(ABC),它的要求远不止“可迭代+可索引”:
- 必须实现
index()、count()方法 __getitem__必须支持切片并返回同类型的序列- 还要满足其他Python原生序列的协议细节
而numpy的ndarray并没有实现这些全部要求(比如没有index()方法),所以Mypy不认可它是Sequence的子类。这和可变/不可变无关——比如list是可变的,但它完全实现了Sequence协议,所以能被Mypy接受。
满足“可迭代+可索引”需求的类型提示方案
方案1:使用联合类型兼容原生序列和numpy数组
直接把参数类型设为Sequence[float]和numpy数组类型的联合,这样两种类型都能通过检查:
import numpy as np from typing import Sequence from numpy.typing import NDArray def compute(x: Sequence[float] | NDArray[np.floating[Any]]) -> bool: # 原计算逻辑不变 ... compute(np.linspace(0, 1, 10)) # 正常通过类型检查 compute([1.0, 2.0, 3.0]) # 原生列表也能正常使用
方案2:自定义通用协议
如果想支持所有“可索引+可迭代+有长度”的对象(不限于Python序列和numpy数组),可以自定义一个Protocol:
import numpy as np from typing import Protocol, TypeVar T = TypeVar('T', covariant=True) class IndexableIterable(Protocol[T]): def __getitem__(self, index: int) -> T: ... def __len__(self) -> int: ... def __iter__(self) -> iter[T]: ... def compute(x: IndexableIterable[float]) -> bool: # 原计算逻辑不变 ... compute(np.linspace(0, 1, 10)) # 正常通过 compute([1.0, 2.0]) # 列表支持 compute((1.0, 2.0)) # 元组也支持
这个协议只定义了你需要的核心行为,不会引入额外的不必要要求,兼容性更强。
方案3:仅针对numpy数组(如果不需要支持原生序列)
如果你的函数只需要处理numpy数组,直接指定numpy的类型即可:
import numpy as np from numpy.typing import NDArray def compute(x: NDArray[np.floating[Any]]) -> bool: ... compute(np.linspace(0, 1, 10))
内容的提问来源于stack exchange,提问作者natty
相关产品推荐
相关产品推荐

