为含NumPy的基础Python函数优化类型提示方案咨询
更简洁的NumPy兼容类型提示方案
替代@overload的简化方案
不用写多个重载签名,直接通过类型别名+联合类型/泛型TypeVar就能实现需求,结合mypy的numpy插件可完美支持类型检查:
1. 基础联合类型方案
先定义一个包含所有合法输入类型的别名,直接用于函数参数和返回值提示:
import numpy as np from numpy.typing import NDArray from typing import Union # 定义合法输入类型:int、float、任意数值类型的NumPy数组 Numeric = Union[int, float, NDArray[np.number]] def my_math_func(x: Numeric) -> Numeric: # 示例实现:简单的数值运算 return x * 2
- 效果:mypy会自动拦截字符串、布尔值等非法输入,同时允许int、float、NumPy数值数组的合法调用
- 说明:
NDArray[np.number]限定数组为数值类型(排除字符串数组等),如果需要允许任意类型数组,可改为NDArray[Any]
2. 严格类型匹配方案(输入输出类型一致)
如果需要函数返回类型与输入类型严格匹配(比如输入int返回int,输入数组返回同类型数组),用TypeVar实现泛型提示:
import numpy as np from numpy.typing import NDArray from typing import TypeVar # 定义泛型类型变量,限定为合法的数值类型 T = TypeVar('T', int, float, NDArray[np.number]) def my_math_func(x: T) -> T: return x * 2
- 效果:mypy会检查返回值类型是否与输入一致,避免出现输入int却返回float的情况
NumPy类型提示示例资源
以下是无需外链的实用参考内容:
- NumPy官方类型提示文档:内置的类型提示章节包含:
- 基础数组类型:
NDArray[np.int32](一维int32数组)、NDArray[np.float64, np.shape[2, 3]](2×3的float64二维数组) - 泛型数组类型的使用方法
- 基础数组类型:
- mypy numpy插件文档:说明如何配置插件支持NumPy类型检查,以及常见错误的解决方式
- NumPy内置函数参考:查看NumPy自带函数的类型提示(比如
np.add、np.mean),可通过源码或help(np.add)获取,作为自定义函数的参考 - 实用代码示例:
# 接受任意数值类型的一维数组,返回同类型数组 def square_array(arr: NDArray[np.number]) -> NDArray[np.number]: return arr ** 2 # 接受int/float或任意维度的float数组,返回float结果 def calculate_mean(x: Union[int, float, NDArray[np.float64]]) -> float: return np.mean(x) if isinstance(x, np.ndarray) else x
内容的提问来源于stack exchange,提问作者j-hil
相关产品推荐
相关产品推荐

