如何为NumPy函数用TypeVar标注输入输出dtype适配Pyright?
NumPy函数输入输出dtype一致性的类型提示问题
尝试为NumPy函数标注输入输出的dtype一致性时,使用TypeVar绑定np.dtype触发了Pyright报错。
示例代码
import numpy as np from numpy.typing import NDArray from typing import TypeVar T = TypeVar("T", bound=np.dtype) A = TypeVar("A") def generic_example(value: A) -> A: print(type(value)) return value def create_from_array(array: NDArray[T]) -> NDArray[T]: return array * np.arange(array.size, dtype=array.dtype).reshape(array.shape) if __name__ == "__main__": array = np.array([[1, 2, 3, 4], [4, 3, 2, 1]]) mylist = generic_example(list(range(19))) array = generic_example(array) print(create_from_array(array))
Pyright报错信息
.../numpy_typing_test.py .../numpy_typing_test.py:14:30 - error: Could not specialize type "NDArray[ScalarType@NDArray]" Type "T@create_from_array" cannot be assigned to type "generic" "dtype[Unknown]*" is incompatible with "generic" .../numpy_typing_test.py:14:45 - error: Could not specialize type "NDArray[ScalarType@NDArray]" Type "T@create_from_array" cannot be assigned to type "generic" "dtype[Unknown]*" is incompatible with "generic" .../numpy_typing_test.py:15:12 - error: Operator "*" not supported for types "NDArray[Unknown]" and "ndarray[Any, dtype[ScalarType@NDArray]]" (reportGeneralTypeIssues) 3 errors, 0 warnings, 0 informations
问题原因
NDArray的泛型参数要求是标量类型(比如int、float32这类Python/NumPy数值类型),而非np.dtype对象。之前绑定np.dtype到TypeVar,导致类型不匹配,触发Pyright报错。
修正后的代码
import numpy as np from numpy.typing import NDArray, ScalarType from typing import TypeVar A = TypeVar("A") # 绑定到NumPy支持的标量类型,而非dtype对象 T = TypeVar("T", bound=ScalarType) def generic_example(value: A) -> A: print(type(value)) return value def create_from_array(array: NDArray[T]) -> NDArray[T]: # array.dtype会被推断为np.dtype[T],与输入数组的标量类型匹配 return array * np.arange(array.size, dtype=array.dtype).reshape(array.shape) if __name__ == "__main__": array = np.array([[1, 2, 3, 4], [4, 3, 2, 1]]) mylist = generic_example(list(range(19))) array = generic_example(array) print(create_from_array(array))
关键说明
- 从
numpy.typing导入ScalarType,它是NumPy所有支持的标量类型的联合类型,作为TypeVar的绑定类型,完全符合NDArray泛型参数的要求。 - 修正后,Pyright能正确推断输入输出数组的dtype一致性:
array.dtype会被识别为与T对应的np.dtype实例,np.arange生成的数组类型也会与输入数组匹配,彻底解决运算符不支持的报错。
内容的提问来源于stack exchange,提问作者vahvero
相关产品推荐
相关产品推荐

