无符号整数Numpy数组的正确类型提示方案问询
Numpy无符号整数数组的类型注解方案
需求
为处理无符号整数Numpy数组的函数添加类型注解:
- 限制输入为无符号整数类型的数组,排除浮点型等无关类型
- 保持输入与输出数组的dtype一致
- 避免过于宽泛(允许所有dtype)或过于具体(仅指定单个dtype)的注解
尝试的错误方案及报错
方案1:直接绑定np.unsignedinteger
from typing import TypeVar, Tuple import numpy as np import numpy.typing as npt # 报错:Missing type parameters for generic type "unsignedinteger" S = TypeVar("S", bound=np.unsignedinteger) def subsample_unsigned_S(x: npt.NDArray[S], grid: Tuple[slice, ...]) -> npt.NDArray[S]: return x[grid]
报错原因:
np.unsignedinteger是泛型类,必须指定类型参数才能作为TypeVar的bound。
方案2:将数组类型作为TypeVar候选值
U = TypeVar( "U", npt.NDArray[np.ubyte], npt.NDArray[np.ushort], npt.NDArray[np.uintc], npt.NDArray[np.uint], npt.NDArray[np.ulonglong], ) # 报错:Type argument "U" of "NDArray" must be a subtype of "generic" def subsample_unsigned_U(x: npt.NDArray[U], grid: Tuple[slice, ...]) -> npt.NDArray[U]: return x[grid]
报错原因:
NDArray的类型参数要求是numpy的dtype类型(如np.ubyte),而不是整个数组类型。
方案3:同时指定候选值和bound
V = TypeVar( "V", npt.NDArray[np.ubyte], npt.NDArray[np.ushort], npt.NDArray[np.uintc], npt.NDArray[np.uint], npt.NDArray[np.ulonglong], bound=np.generic, ) # 报错:TypeVar cannot have both values and an upper bound def subsample_unsigned_V(x: npt.NDArray[V], grid: Tuple[slice, ...]) -> npt.NDArray[V]: return x[grid]
报错原因:TypeVar不允许同时设置候选值列表和upper bound。
方案4:Union定义数组类型
from typing import Union UnsignedIntegerArray = Union[ npt.NDArray[np.ubyte], npt.NDArray[np.ushort], npt.NDArray[np.uintc], npt.NDArray[np.uint], npt.NDArray[np.ulonglong], ] def subsample_unsigned(x: UnsignedIntegerArray, grid: Tuple[slice, ...]) -> UnsignedIntegerArray: return x[grid]
问题:无法关联输入与输出的具体dtype,mypy会认为返回值可能是任意无符号数组类型,而非与输入一致。
正确解决方案
方案A:绑定泛型无符号整数类型
通过给np.unsignedinteger指定泛型参数np.generic,解决方案1的报错问题,同时限制所有无符号整数dtype:
from typing import TypeVar, Tuple import numpy as np import numpy.typing as npt # 绑定到所有无符号整数dtype UIntDType = TypeVar("UIntDType", bound=np.unsignedinteger[np.generic]) def subsample_unsigned( x: npt.NDArray[UIntDType], grid: Tuple[slice, ...] ) -> npt.NDArray[UIntDType]: """仅接收无符号整数数组,返回同dtype的数组""" return x[grid]
这种写法允许所有无符号整数类型(包括自定义的无符号dtype),同时保证输入输出类型一致。
方案B:枚举具体无符号dtype
如果需要严格限制为numpy内置的无符号整数类型,可以直接将这些dtype作为TypeVar的候选值:
from typing import TypeVar, Tuple import numpy as np import numpy.typing as npt # 仅允许numpy内置的无符号整数dtype UIntDType = TypeVar( "UIntDType", np.ubyte, np.ushort, np.uintc, np.uint, np.ulonglong, ) def subsample_unsigned( x: npt.NDArray[UIntDType], grid: Tuple[slice, ...] ) -> npt.NDArray[UIntDType]: return x[grid]
这种写法更严格,只接受指定的几种无符号类型,同时保持输入输出类型关联。
内容的提问来源于stack exchange,提问作者Niklas Netter
相关产品推荐
相关产品推荐

