NumPy 2.5 类型提示增强:numpy.linalg 形状推断(shape-typing)全面解析
【免费下载链接】numpyThe fundamental package for scientific computing with Python.项目地址: https://gitcode.com/gh_mirrors/nu/numpy
本篇技术指南聚焦 NumPy 在类型提示(type hints)领域的一项实质性改进——numpy.linalg各函数现在会根据输入数组的形状类型(shape-type)推断返回数组的形状与 dtype 类型。文中将结合本仓库中发布说明、类型桩(stub)源码与类型测试,完整还原其实现机制、适用边界与验证方法,帮助你准确理解reveal_type输出从ndarray[tuple[Any, ...], ...]收敛为ndarray[tuple[int], dtype[float64]]这一变化背后的原理。
变更概览:linalg 返回类型从“任意形状”走向“精确形状”
本仓库的发布说明 doc/release/upcoming_changes/32424.typing.rst 记录了一项面向类型检查器的改进:numpy.linalg中的函数现在会使用其输入的形状类型来推断返回数组的形状类型。也就是说,当静态类型检查器(如 mypy、pyright)能够确定输入是几维数组时,它也能顺带确定返回值是几维数组,而不再统一退化为tuple[Any, ...]这样的“任意形状”。
以np.linalg.solve为例,发布说明给出的对比是:
A = np.array([[1.1, 1.2], [2.1, 2.2]]) b = np.array([1.0, 2.0]) reveal_type(np.linalg.solve(A, b)) # before: ndarray[tuple[Any, ...], dtype[float64]] # after: ndarray[tuple[int], dtype[float64]]再如np.linalg.inv:
stacked = np.array([[[1.0, 0.0], [0.0, 1.0]]]) reveal_type(np.linalg.inv(stacked)) # before: ndarray[tuple[Any, ...], dtype[Any]] # after: ndarray[tuple[int, int, int], dtype[float64]]inv的例子同时展示了两层改进:形状维度从tuple[Any, ...]精确到了具体的tuple[int, int, int],dtype 也从dtype[Any]收敛到了dtype[float64]。
背景:NumPy 类型模型中 ndarray 的两个泛型参数
要理解本次改进,需要先了解 NumPy 类型桩中np.ndarray的泛型设计。在 NumPy 的类型模型中,np.ndarray携带两个泛型参数:形状(shape)与 dtype。形状参数的核心类型别名定义在 numpy/_typing/_shape.py 中:
type _Shape = tuple[int, ...] type _AnyShape = tuple[Any, ...] # Anything that can be coerced to a shape tuple type _ShapeLike = SupportsIndex | Sequence[SupportsIndex]其中:
_Shape(tuple[int, ...])代表“各维度都是已知整数”的精确形状,例如tuple[int, int]表示一个二维数组;_AnyShape(tuple[Any, ...])代表“完全未知”的任意形状,也是此前 linalg 返回类型的常见形态。
本次改进的本质,就是让 linalg 的重载(overload)返回值尽量使用_Shape约束下的具体形状类型,而不是_AnyShape。从本仓库的变更记录 doc/changelog/2.5.0-changelog.rst 可以看到,这是一系列渐进工作的成果:linalg.diagonal/trace、linalg.lstsq、linalg.outer、linalg.*norm、linalg.cross、linalg.matmul等函数在 2.5.0 版本周期内陆续完成了 shape-typing 与 dtype 特化。
源码级实现:linalg 类型桩中的形状别名与重载体系
linalg 的全部类型桩位于 numpy/linalg/_linalg.pyi,并由 numpy/linalg/init.pyi 对外重导出(re-export)。
形状别名与维度约束
文件开头定义了从一维到四维的精确形状别名,以及“至少 N 维”的约束类型:
type _1D = tuple[int] type _2D = tuple[int, int] type _3D = tuple[int, int, int] type _4D = tuple[int, int, int, int] type _AtMost1D = tuple[()] | _1D type _AtLeast1D = tuple[int, *tuple[int, ...]] type _AtLeast2D = tuple[int, int, *tuple[int, ...]] type _AtLeast3D = tuple[int, int, int, *tuple[int, ...]] type _AtLeast4D = tuple[int, int, int, int, *tuple[int, ...]]在此基础上定义了形状精确的数组类型:
type _Array0D[ScalarT: np.generic] = np.ndarray[tuple[()], np.dtype[ScalarT]] type _Array1D[ScalarT: np.generic] = np.ndarray[_1D, np.dtype[ScalarT]] type _Array2D[ScalarT: np.generic] = np.ndarray[_2D, np.dtype[ScalarT]] type _Array3D[ScalarT: np.generic] = np.ndarray[_3D, np.dtype[ScalarT]] type _Array4D[ScalarT: np.generic] = np.ndarray[_4D, np.dtype[ScalarT]]这些别名是重载返回值形状推断的直接载体:只要匹配到“二维输入”重载,返回值类型便直接写成_Array2D[...],形状即固定为tuple[int, int]。
输入侧的形状感知协议
为了让“已知形状的 ndarray 输入”能被重载精确匹配,类型桩还定义了一个形状泛型版本的__array__协议:
@type_check_only class _SupportsArrayShapeT: _Shape, DTypeT: np.dtype: def __array__(self, /) -> np.ndarray[ShapeT, DTypeT]: ...配合输入侧别名(如_ArrayLike2D2、_ArrayLike3D2、_Sequence2D、_Sequence3D),类型检查器可以在“传入的是二维 ndarray / 二维嵌套序列”时选中对应的低秩重载。
solve 的维度组合重载
solve是发布说明中最典型的例子。在 numpy/linalg/_linalg.pyi 中,solve(a, b)针对输入形状组合的每一种已知情形都有独立重载,这里摘录核心的几组:
# 2d + 1d -> 1d def solve(a: _ArrayLike2D2[_to_float64, float], b: _ArrayLike1D2[_to_float64, float]) -> _Array1D[np.float64]: ... # 2d + 2d -> 2d def solve(a: _ArrayLike2D2[_to_float64, float], b: _ArrayLike2D2[_to_float64, float]) -> _Array2D[np.float64]: ... # 2d + 3d -> 3d def solve(a: _ArrayLike2D2[_to_float64, float], b: _ArrayLike3D2[_to_float64, float]) -> _Array3D[np.float64]: ... # 3d + 1d -> 2d def solve(a: _ArrayLike3D2[_to_float64, float], b: _ArrayLike1D2[_to_float64, float]) -> _Array2D[np.float64]: ... # 3d + 2d -> 3d def solve(a: _ArrayLike3D2[_to_float64, float], b: _ArrayLike2D2[_to_float64, float]) -> _Array3D[np.float64]: ...这与solve的运行时语义完全一致:求解Ax = b时,结果x的维度等于b的维度(b为一维向量则返回一维,b为矩阵则返回矩阵)。这正是发布说明中solve(A, b)返回ndarray[tuple[int], dtype[float64]]的桩层面来源。
在形状未知的兜底重载(fallback)中,返回值才退化为NDArray[np.float64]甚至NDArray[Any]:
# fallback def solve(a: _ArrayLikeComplex_co, b: _ArrayLikeComplex_co) -> NDArray[Any]: ...inv 的“已知数组原样返回”重载
inv的桩设计采用了另一条路径:如果传入的是“已知形状、已知 dtype 的 ndarray”,则直接保留输入的形状与 dtype 泛型参数返回:
@overload # known array def inv[ArrayT: np.ndarray[_AtLeast2D, np.dtype[_inexact32 | _inexact64]]](a: ArrayT) -> ArrayT: ... @overload # known shape, known dtype def inv[ShapeT: _AtLeast2D, DTypeT: np.dtype[_inexact32 | _inexact64]]( a: _SupportsArray[ShapeT, DTypeT], ) -> np.ndarray[ShapeT, DTypeT]: ... @overload # 2d +float def inv(a: Sequence[Sequence[float | np.integer | np.bool]]) -> _Array2D[np.float64]: ... @overload # 3d +float def inv(a: Sequence[Sequence[Sequence[float | np.integer | np.bool]]]) -> _Array3D[np.float64]: ...于是发布说明中的stacked = np.array([[[1.0, 0.0], [0.0, 1.0]]])作为三维数组传入时,匹配“2d/3d 已知形状”重载,返回值形状即为tuple[int, int, int],dtype 因浮点输入被推定为float64,两者共同构成ndarray[tuple[int, int, int], dtype[float64]]。
更多函数的形状语义
同类模式扩展到了其他 linalg 函数,可结合类型桩逐一对照:
det:二维输入返回标量np.float64/np.complex128,三维及以上的批量输入返回一维数组(_Array1D[np.float64]),见 numpy/linalg/_linalg.pyi 中det的重载组;qr:返回QRResult[InexactT, ShapeT]泛型 NamedTuple,Q、R 均为形状泛型数组;eig/eigh:EigResult、EighResult携带形状与 dtype 两个泛型参数,特征值/特征向量的维度得以精确表达;svd:SVDResult泛型化 U、S、Vh 三者的形状与 dtype;lstsq:_LstSqResult[ShapeT, InexactT, FloatingT]携带解的精确形状;norm/matrix_rank/pinv/tensorsolve/tensorinv等同样具备分维度的重载组。
需要指出,cholesky、tensorinv等函数的桩注释中带有 “keep in sync with the other inverse functions and cholesky” 的同步维护提示,说明这些重载组在开发中被刻意保持结构一致,以便统一演进。
适用范围与已知限制
发布说明中明确强调了一句关键限定:由于 Python 类型系统的限制,这种形状推断通常只适用于低秩(low-rank)的类数组对象。
原因可以从桩代码中直接读出:
- 重载组合随维度爆炸:
solve需要为 1d/2d/3d 的每种组合单独书写重载,维度越高组合越多;类型桩目前覆盖到三维组合(_Array3D),更高的秩没有逐一枚举,落入_AtLeast*D或兜底重载后形状信息自然丢失。 - 嵌套序列的类型化成本高:要推断
list[list[float]]是“二维”,类型检查器必须逐层分析序列嵌套深度;当输入是形状未知的通用ArrayLike(例如np.ndarray不带形状泛型参数)时,只能匹配 fallback 重载。 - 形状是“类型”而非“值”:
_Shape中的每个元素都是int类型而非具体数值,因此类型系统能表达“返回值是一维/二维/三维”,但无法表达“长度为 n 的方阵”这类具体尺寸约束。
因此,在实际项目中,只有当输入是形状可静态确定的 ndarray 或嵌套序列时,返回值形状推断才会生效;对于从函数参数一路传进来的“未知形状数组”,类型仍会退化为NDArray[Any]形态,这属于预期行为而非回归。
测试验证:reveal 测试如何保证推断正确
本次改进并非无测试保障。仓库的 typing 测试体系通过 mypy 的 reveal 机制逐条校验推断结果:
- 测试用例数据位于 numpy/typing/tests/data/reveal/linalg.pyi,其中使用
assert_type对每个重载分支的返回类型做精确断言。例如:
assert_type(np.linalg.solve(AR_f8_2d, AR_f8_1d), _Array1D[np.float64]) assert_type(np.linalg.solve(AR_f8_2d, AR_f8_2d), _Array2D[np.float64]) assert_type(np.linalg.solve(AR_f8_2d, AR_f8_3d), _Array3D[np.float64]) assert_type(np.linalg.solve(AR_f8_3d, AR_f8_1d), _Array2D[np.float64]) assert_type(np.linalg.inv(AR_f8_2d), _Array2D[np.float64]) assert_type(np.linalg.inv(AR_f8_3d), _Array3D[np.float64]) assert_type(np.linalg.qr(AR_f8_2d), QRResult[np.float64, _2D])其中AR_f8_2d、AR_f8_3d等变量在文件头部被定义为形状精确的数组类型(如_Array2D[np.float64]),从而精确驱动重载选择。
- 测试驱动逻辑在 numpy/typing/tests/test_typing.py 的
test_reveal中:它读取data/reveal目录下的用例文件,运行 mypy 收集每条表达式的推断输出,并逐一与预期比对,任何一条推断与桩不符都会导致测试失败。
这意味着,如果你在自己的项目中观察到np.linalg.solve的推断结果与本文示例不一致,首先应确认输入对象是否为“形状可静态确定”的 ndarray 或嵌套序列——这正是 reveal 用例所覆盖的输入形态。
对使用者的影响与建议
本次变更属于纯类型层面改进,不改变任何运行时行为:np.linalg.solve、np.linalg.inv等函数的数值计算语义、参数与返回值与之前完全一致,升级 NumPy 不会影响既有代码的运行结果。它影响的是静态类型检查器的推断精度:
- 更早发现维度错误:此前
solve的返回值类型是ndarray[tuple[Any, ...], ...],将其误当作三维数组使用不会触发类型告警;改进后返回值明确为_Array1D/_Array2D/_Array3D,把一维结果当矩阵索引时,mypy/pyright 会给出提示。 - IDE 补全与文档体验提升:返回类型的形状明确后,编辑器对后续链式调用的自动补全(如
.shape、索引运算后的类型)会更为准确。 - dtype 透明度提升:如
inv示例所示,整型/浮点输入现在能收敛到dtype[float64],复数输入收敛到dtype[complex128],减少了dtype[Any]对下游类型推导的污染。
若希望在自己的代码中获得最佳推断效果,建议:为函数参数显式标注形状明确的 ndarray 类型(例如npt.NDArray[np.float64]至少保留 dtype 信息);在确实需要跨维度通用时,接受 fallback 重载给出的宽松类型,而不是依赖对Any的隐式假设。
小结
numpy.linalg的 shape-typing 改进,是 NumPy 类型系统从“只关心 dtype”走向“形状与 dtype 并重”的又一步:通过 numpy/linalg/_linalg.pyi 中成体系的形状别名、_SupportsArray协议与按维度组合书写的重载,静态类型检查器得以在低秩输入场景下精确推断返回值形状;numpy/typing/tests/data/reveal/linalg.pyi 中的assert_type断言则保证了这套推断的持续正确性。对于日常使用solve、inv、det、qr、eig等线性代数例程的开发者而言,这意味着类型检查从“放行一切”变成了“帮助纠错”,且完全无需改动任何运行时代码。
【免费下载链接】numpyThe fundamental package for scientific computing with Python.项目地址: https://gitcode.com/gh_mirrors/nu/numpy
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考