如何为Numpy NDArray提供整数形式的SHAPE类型提示?
解决Numpy NDArray用整数定义固定形状类型提示的问题
如果你想在新版Numpy里用整数直接指定NDArray的固定形状(比如3行5列的浮点数组),可以用下面两种方法:
方法一:使用Literal字面量类型
把指定形状的整数包装成Literal类型,让类型检查器识别这是固定值:
import numpy as np from numpy.typing import NDArray from typing import Literal # 直接内联字面量 def func(x: NDArray[(Literal[3], Literal[5]), float]): pass # 或者先定义变量再使用 rows = Literal[3] columns = Literal[5] def func2(x: NDArray[(rows, columns), float]): pass
方法二:使用Final常量
用Final修饰形状变量,标记它是不可变的固定值,Numpy的类型系统会认可这种用法:
import numpy as np from numpy.typing import NDArray from typing import Final rows: Final[int] = 3 columns: Final[int] = 5 def func(x: NDArray[(rows, columns), float]): pass
为什么原来的写法不行?
新版Numpy对NDArray的形状参数做了约束,要求形状必须是字面量值、字符串或者能被类型系统识别为固定值的表达式。普通变量因为无法确定是否会被修改,所以不被允许,而Literal和Final能明确告诉类型系统这是固定不变的整数,因此可以合法用作形状参数。
内容的提问来源于stack exchange,提问作者TigaJi
相关产品推荐
相关产品推荐

