如何为支持任意NumPy浮点类型的数组添加类型提示?
如何约束NDArray接受任意NumPy浮点类型
当然可以,你不用宽泛的NDArray[Any],而是用NumPy类型系统提供的专用别名来精准约束,确保只接受np.floating的子类类型。
具体实现代码
from numpy.typing import NDArray, floating import numpy as np # 定义仅接受任意NumPy浮点类型的数组变量 any_np_float: NDArray[floating] # 测试用例:所有浮点类型都会通过类型检查 arr_float32 = np.array([1.2, 3.4], dtype=np.float32) arr_float64 = np.array([5.6, 7.8], dtype=np.float64) arr_float16 = np.array([9.0, 0.1], dtype=np.float16) any_np_float = arr_float32 # 合法 any_np_float = arr_float64 # 合法 any_np_float = arr_float16 # 合法 # 非浮点类型会触发类型检查错误 arr_int = np.array([1, 2], dtype=np.int32) any_np_float = arr_int # 类型不匹配,静态检查工具会报错
关键说明
numpy.typing.floating是NumPy官方提供的类型别名,它涵盖了所有np.floating的子类(包括np.float16、np.float32、np.float64等常见浮点类型)。- 相比
NDArray[Any],这种写法既保留了对所有浮点类型的兼容性,又能通过静态类型检查排除整数、布尔值等非浮点类型的数组,让类型定义更严谨。
内容的提问来源于stack exchange,提问作者ilovewt
相关产品推荐
相关产品推荐

