You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 17:42:25