类型注解 (numpy.typing)#
新功能,版本 1.20。
NumPy API 的大部分内容都具有符合 PEP 484 标准的类型注解。此外,还为用户提供了一些类型别名,最主要的是下面两个:
Mypy 插件#
1.21 版本新功能。
一个用于管理平台特定注解的 mypy 插件。其功能可分为三个不同的部分:
分配特定
number子类的(平台相关)精度,包括int_、intp和longlong等。有关受影响类的全面概述,请参阅关于 标量类型 的文档。如果没有该插件,所有相关类的精度都将被推断为Any。移除该平台不可用的所有扩展精度
number子类。最显著的是包括float128和complex256。如果没有该插件,在 mypy 看来,所有扩展精度类型在所有平台上都是可用的。分配
c_intp的(平台相关)精度。如果没有该插件,该类型将默认指向ctypes.c_int64。新功能,版本 1.22。
2.3 版本弃用:numpy.typing.mypy_plugin 入口点已被弃用,建议改用平台无关的静态类型推断。请从 mypy 配置的 plugins 部分移除 numpy.typing.mypy_plugin;如果这样做引发了新的错误,请提交一个包含最小重现代码的问题。
示例#
要启用该插件,必须将其添加到 mypy 的 配置文件中。
[mypy]
plugins = numpy.typing.mypy_plugin
与运行时 NumPy API 的差异#
NumPy 非常灵活。试图静态描述所有可能性会导致类型定义变得不够实用。因此,类型化的 NumPy API 通常比运行时的 NumPy API 更严格。本节描述了一些显著差异。
ArrayLike#
ArrayLike 类型试图避免创建对象数组。例如:
>>> np.array(x**2 for x in range(10))
array(<generator object <genexpr> at ...>, dtype=object)
这是合法的 NumPy 代码,会创建一个 0 维对象数组。但在使用 NumPy 类型时,类型检查器会针对上述示例发出警告。如果您确实打算这样做,可以使用 # type: ignore 注释
>>> np.array(x**2 for x in range(10)) # type: ignore
或者显式地将类数组对象标记为 Any
>>> from typing import Any
>>> array_like: Any = (x**2 for x in range(10))
>>> np.array(array_like)
array(<generator object <genexpr> at ...>, dtype=object)
DTypeLike#
DTypeLike 类型试图避免使用如下所示的字段字典来创建 dtype 对象:
>>> x = np.dtype({"field1": (float, 1), "field2": (int, 3)})
虽然这是合法的 NumPy 代码,但类型检查器会发出警告,因为不建议这样使用。请参阅:数据类型对象
数值精度#
numpy.number 子类的精度被视为不变的泛型参数(请参阅 NBitBase),这简化了涉及基于精度的类型转换过程的注解。
>>> from typing import TypeVar
>>> import numpy as np
>>> import numpy.typing as npt
>>> T = TypeVar("T", bound=npt.NBitBase)
>>> def func(a: np.floating[T], b: np.floating[T]) -> np.floating[T]:
... ...
因此,float16、float32 和 float64 等依然是 floating 的子类型,但与运行时不同,它们不一定被视为子类。
2.3 版本弃用:NBitBase 辅助工具已被弃用,将在未来的版本中移除。建议通过 typing.overload 或以具体标量类为上界的 TypeVar 定义来表达精度关系。例如:
from typing import TypeVar import numpy as np S = TypeVar("S", bound=np.floating) def func(a: S, b: S) -> S: ...或者在不同的输入类型映射到不同的输出类型时:
from typing import overload
import numpy as np
@overload
def phase(x: np.complex64) -> np.float32: ...
@overload
def phase(x: np.complex128) -> np.float64: ...
@overload
def phase(x: np.clongdouble) -> np.longdouble: ...
def phase(x: np.complexfloating) -> np.floating:
...
Timedelta64#
timedelta64 类在静态类型检查期间不被视为 signedinteger 的子类,前者仅继承自 generic。
记录数组数据类型 (Record array dtypes)#
numpy.recarray 的 dtype 以及通用的 创建记录数组 函数可以通过以下两种方式之一指定:
直接通过
dtype参数。使用最多五个辅助参数,它们通过
numpy.rec.format_parser进行操作:formats、names、titles、aligned和byteorder。
这两种方法目前被标记为互斥,即如果指定了 dtype,则不能指定 formats。虽然这种互斥性在运行时并没有(严格)强制执行,但结合使用两种 dtype 指定器可能会导致意外甚至完全错误的行为。
API#
ndarray
numpy.ndarray 类是一个接受两个类型参数的 泛型类型:
numpy.ndarray.shape的类型,必须是 int 的 tuple,例如tuple[int, int](2-D 形状) 或tuple[()](0-D 形状)。默认形状为tuple[Any, ...],表示未知形状及任意维数。目前不支持Literal整数或其他更具体的类型。numpy.ndarray.dtype的类型,必须是numpy.dtype的子类型,例如numpy.dtype[numpy.float64]。如果省略,则默认指向numpy.dtype[Any]。
>>> import numpy as np
>>> type ImageRGB = np.ndarray[tuple[int, int, int], np.dtype[np.uint8]]
>>> type Vector[S: np.generic] = np.ndarray[tuple[int], np.dtype[S]]
- numpy.typing.ArrayLike = typing.Union[...]#
一个代表可强制转换为
ndarray的对象的Union。其中包括:
标量。
(嵌套)序列。
实现 __array__ 协议的对象。
新功能,版本 1.20。
另请参阅
- array_like:
任何可解释为 ndarray 的标量或序列。
示例
>>> import numpy as np >>> import numpy.typing as npt >>> def as_array(a: npt.ArrayLike) -> np.ndarray: ... return np.array(a)
- numpy.typing.DTypeLike = typing.Union[...]#
-
其中包括:
新功能,版本 1.20。
另请参阅
- 指定和构造数据类型
所有可强制转换为数据类型的对象的全面概述。
示例
>>> import numpy as np >>> import numpy.typing as npt >>> def as_dtype(d: npt.DTypeLike) -> np.dtype: ... return np.dtype(d)
- numpy.typing.NDArray = NDArray#
一个
np.ndarray[tuple[Any, ...], np.dtype[ScalarT]]类型别名,就其dtype.type而言是 泛型 的。可在运行时用于对具有给定 dtype 和未指定形状的数组进行类型标注。
1.21 版本新功能。
示例
>>> import numpy as np >>> import numpy.typing as npt >>> print(npt.NDArray) NDArray >>> print(npt.NDArray[np.float64]) NDArray[numpy.float64] >>> NDArrayInt = npt.NDArray[np.int_] >>> a: NDArrayInt = np.arange(10) >>> def func(a: npt.ArrayLike) -> npt.NDArray[Any]: ... return np.array(a)
- class numpy.typing.NBitBase[source]#
在静态类型检查期间代表
numpy.number精度的类型。NBitBase仅用于静态类型检查,代表分层子类集的基类。随后的每个子类都在此用于表示较低级别的精度,例如64Bit > 32Bit > 16Bit。新功能,版本 1.20。
2.3 版本弃用:请改用
@typing.overload或以标量类型为上界的TypeVar。示例
以下是一个典型的用法示例:
NBitBase在此用于注解一个函数,该函数接收任意精度的浮点数和整数作为参数,并返回精度较大者的新浮点数(例如np.float16 + np.int64 -> np.float64)。>>> from typing import TYPE_CHECKING >>> import numpy as np >>> import numpy.typing as npt >>> def add[S: npt.NBitBase, T: npt.NBitBase]( ... a: np.floating[S], b: np.integer[T] ... ) -> np.floating[S | T]: ... return a + b >>> a = np.float16() >>> b = np.int64() >>> out = add(a, b) >>> if TYPE_CHECKING: ... reveal_locals() ... # note: Revealed local types are: ... # note: a: numpy.floating[numpy.typing._16Bit*] ... # note: b: numpy.signedinteger[numpy.typing._64Bit*] ... # note: out: numpy.floating[numpy.typing._64Bit*]