如何为接收泛型浮点数组的numpy函数编写类型标注
问题原因
你遇到的TypeError: 'type' object is not subscriptable报错,是因为运行时Python解释器会执行标注里的下标语法(也就是[]),但对应版本的numpy类型在运行时本身不支持下标操作,和静态类型检查阶段的逻辑冲突。
解决方案
方案1:添加future注解声明(兼容性最强,Python 3.7+通用)
在代码最顶部加入from __future__ import annotations,该声明会让Python运行时把所有类型标注自动转为字符串,不会尝试解析下标语法,仅在静态类型检查(比如mypy、pyright)阶段才会被识别解析。
完整可运行代码如下:
from __future__ import annotations from typing import TypeVar, cast import numpy as np import numpy.typing as npt T = TypeVar("T", bound=npt.NBitBase) def entropy(x: npt.NDArray[np.floating[T]]) -> np.floating[T]: return cast(np.floating[T], np.sum(x * np.log(x)))
方案2:手动用引号包裹类型标注
如果不想加future声明,也可以手动把涉及泛型下标、运行时不支持的标注用引号包裹,效果和方案1一致:
from typing import TypeVar, cast import numpy as np import numpy.typing as npt T = TypeVar("T", bound=npt.NBitBase) def entropy(x: "npt.NDArray[np.floating[T]]") -> "np.floating[T]": return cast(np.floating[T], np.sum(x * np.log(x)))
额外注意事项
- 需保证numpy版本 >= 1.21,该版本开始才完整支持
numpy.typing下的泛型类型标注 - 如果不需要严格绑定浮点位数,也可以简化标注为
def entropy(x: npt.NDArray[np.floating]) -> np.floating:,不需要额外声明TypeVar即可满足大部分类型检查场景
内容的提问来源于stack exchange,提问作者shaunc
相关产品推荐
相关产品推荐

