类型注解 (numpy.typing)#

新功能,版本 1.20。

NumPy API 的大部分内容都具有符合 PEP 484 标准的类型注解。此外,还为用户提供了一些类型别名,最主要的是下面两个:

  • ArrayLike: 可转换为数组的对象

  • DTypeLike: 可转换为数据类型(dtype)的对象

Mypy 插件#

1.21 版本新功能。

一个用于管理平台特定注解的 mypy 插件。其功能可分为三个不同的部分:

  • 分配特定 number 子类的(平台相关)精度,包括 int_intplonglong 等。有关受影响类的全面概述,请参阅关于 标量类型 的文档。如果没有该插件,所有相关类的精度都将被推断为 Any

  • 移除该平台不可用的所有扩展精度 number 子类。最显著的是包括 float128complex256。如果没有该插件,在 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]:
...     ...

因此,float16float32float64 等依然是 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 进行操作:formatsnamestitlesalignedbyteorder

这两种方法目前被标记为互斥,如果指定了 dtype,则不能指定 formats。虽然这种互斥性在运行时并没有(严格)强制执行,但结合使用两种 dtype 指定器可能会导致意外甚至完全错误的行为。

API#

ndarray

numpy.ndarray 类是一个接受两个类型参数的 泛型类型

  1. numpy.ndarray.shape 的类型,必须是 inttuple,例如 tuple[int, int] (2-D 形状) 或 tuple[()] (0-D 形状)。默认形状为 tuple[Any, ...],表示未知形状及任意维数。目前不支持 Literal 整数或其他更具体的类型。

  2. 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[...]#

一个代表可强制转换为 dtype 的对象的 Union

其中包括:

  • type 对象。

  • 字符代码或 type 对象的名称。

  • 具有 .dtype 属性的对象。

新功能,版本 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*]