如何为支持数值或类数组输入的NumPy操作函数正确添加类型提示?
Python类型提示:处理数值与NumPy数组兼容的最佳实践
问题背景
当函数需要同时接受普通数值(int/float)和NumPy数组作为参数时,直接标注np.ndarray会导致类型检查工具(如mypy)报错,因为普通数值类型不匹配ndarray类型。例如:
import numpy as np def square(x: np.ndarray): return x**2 num = 9 print(square(num)) # mypy报错:int类型与ndarray不兼容 arr = np.array([3, 6.31, 9, 8.73]) print(square(arr))
mypy错误提示:
error: Argument 1 to "square" has incompatible type "int"; expected "ndarray[Any, Any]" [arg-type]
你尝试的临时方案通过Union[float, np.ndarray[Any, Any]]让mypy通过,但存在不够严谨的问题(比如未显式支持int类型,且Any会丢失数组的维度和dtype信息)。
最佳实践方案
1. 使用精确的Union类型配合numpy官方类型别名
从numpy.typing导入NDArray(替代直接用np.ndarray),结合具体数值类型构建Union,同时显式支持int、float:
import numpy as np from numpy.typing import NDArray from typing import Union NumOrArray = Union[int, float, NDArray[np.number]] def square(x: NumOrArray) -> NumOrArray: return x**2 # 测试代码 num = 9 print(square(num)) # 类型检查通过 arr = np.array([3, 6.31, 9, 8.73]) print(square(arr)) # 类型检查通过
优势:
NDArray[np.number]涵盖所有数值类型的NumPy数组,比NDArray[Any, Any]更精确- 显式包含int、float,避免类型检查遗漏
- 同时标注返回值类型,让类型信息更完整
2. 使用抽象数值类型简化Union
如果不需要区分int和float,可以用numbers.Real抽象类型来涵盖所有实数类型,配合NDArray:
import numpy as np from numpy.typing import NDArray from typing import Union import numbers NumOrArray = Union[numbers.Real, NDArray[np.number]] def square(x: NumOrArray) -> NumOrArray: return x**2
优势:
- 用抽象类型减少重复,
numbers.Real包含int、float、Decimal等所有实数类型 - 保持类型检查的严谨性
3. 针对数组维度/dtype的精确标注
如果需要更严格的类型约束(比如只接受一维浮点数组),可以指定NDArray的参数:
import numpy as np from numpy.typing import NDArray from typing import Union NumOr1DFloatArray = Union[int, float, NDArray[np.float64]] def square(x: NumOr1DFloatArray) -> NumOr1DFloatArray: return x**2
4. 类型守卫(针对需要分支处理的场景)
如果函数内部需要对数值和数组做不同逻辑处理,可以用类型守卫来明确区分类型:
import numpy as np from numpy.typing import NDArray from typing import Union, TypeGuard def is_ndarray(x: Union[int, float, NDArray]) -> TypeGuard[NDArray]: return isinstance(x, np.ndarray) def square(x: Union[int, float, NDArray]) -> Union[int, float, NDArray]: if is_ndarray(x): # 此处x会被类型检查器识别为NDArray return x**2 + np.ones_like(x) else: # 此处x会被识别为int/float return x**2
优势:
- 让类型检查器能准确推断分支内的变量类型,避免类型错误
- 适合复杂逻辑的函数
总结
你的临时方案思路是对的,但可以通过以下方式优化:
- 用
numpy.typing.NDArray替代np.ndarray[Any, Any],获得更精确的数组类型信息 - 显式包含int类型(或用
numbers.Real简化) - 补充返回值的类型提示
- 复杂场景下使用类型守卫
内容的提问来源于stack exchange,提问作者crabulus_maximus
相关产品推荐
相关产品推荐

