如何为numpy的ndarray添加指定长度与元素类型的泛型类型标注?
解决方法
可行,你可以通过给ndarray添加泛型类型参数来指定元素类型和固定长度,具体实现如下:
错误原因
mypy提示"Missing type parameters for generic type 'ndarray'",是因为新版numpy的ndarray是泛型类型,必须明确指定元素类型和形状信息。
修改后的代码
import numpy as np from numpy import ndarray def foo() -> ndarray[np.float64, np.shape[3]]: # 显式指定dtype确保元素为浮点数 a = np.array([1.0, 2.2, 3.3333], dtype=np.float64) print(a) print(type(a)) return a
说明
np.float64指定了数组元素的类型为64位浮点数,你也可以根据需求用np.float32等其他浮点数类型。np.shape[3]明确了数组是长度为3的一维数组,若你返回的数组长度不符合3,mypy会触发类型检查错误。- 显式添加
dtype=np.float64能避免numpy自动推断类型带来的不确定性,让mypy更准确地进行类型校验。
内容的提问来源于stack exchange,提问作者sixtyfootersdude
相关产品推荐
相关产品推荐

