如何为支持多种数据类型的二维NumPy数组添加类型注解?
解决二维数组多类型支持的类型注解问题
针对你需要为支持int、unsigned int、float、double的二维数组方法添加类型注解的需求,以下是可行的解决方案:
方案一:为每个类型的NDArray使用Union
直接将所有允许的二维数组类型通过Union组合,这是最直观且类型检查器能正确识别的方式:
from typing import NoReturn, Union from nptyping import NDArray, Shape, Int, UInt, Float32, Float64 def foo(arr: Union[ NDArray[Shape["*, *"], Int], NDArray[Shape["*, *"], UInt], NDArray[Shape["*, *"], Float32], NDArray[Shape["*, *"], Float64] ]) -> NoReturn: pass
方案二:使用TypeVar简化重复代码
如果需要在多个地方复用这个类型组合,可以用TypeVar定义一个受限的类型变量,绑定到允许的数值类型上:
from typing import NoReturn, TypeVar from nptyping import NDArray, Shape, Int, UInt, Float32, Float64 # 定义仅允许指定类型的TypeVar NumericArrayType = TypeVar("NumericArrayType", Int, UInt, Float32, Float64) def foo(arr: NDArray[Shape["*, *"], NumericArrayType]) -> NoReturn: pass
为什么原代码无法工作
你之前的写法将Union放在NDArray的第二个类型参数位置,而nptyping.NDArray的第二个参数期望接收单一的nptyping数值类型(如Int、UInt),直接传入Union会导致类型检查器无法解析,因此报错。
额外方案:使用numpy原生类型注解
如果你愿意切换到numpy官方的类型注解方案,可以这样写:
import numpy as np from numpy.typing import NDArray from typing import NoReturn, Union def foo(arr: NDArray[Union[np.int_, np.uint_, np.float32, np.float64]]) -> NoReturn: pass
内容的提问来源于stack exchange,提问作者soumeng78
相关产品推荐
相关产品推荐

