Numpy ndarray子类__getitem__方法的类型标注问题
问题原因与解决方案
核心原因
numpy的ndarray.__getitem__在官方类型标注中,返回值被定义为宽泛的ndarray[Any, Any]类型。虽然运行时numpy会通过__array_finalize__等机制自动返回自定义子类实例,但静态类型检查器(如pyright)无法识别这种动态运行逻辑,只能依赖静态类型签名判断,因此会默认切片操作返回基础ndarray类型,而非你的ArraySubClass。
这既不是pyright的bug,也不算numpy类型签名的错误——numpy的类型标注是为覆盖通用场景设计,无法提前适配所有自定义子类。
解决方案:重写__getitem__方法
要让pyright正确识别切片返回的是ArraySubClass,需在子类中显式重写__getitem__并指定返回类型。示例代码如下:
import numpy as np from typing import Union, Tuple, Any class ArraySubClass(np.ndarray): def __new__(cls, input_array): obj = np.asarray(input_array).view(cls) return obj def __getitem__(self, key: Union[int, slice, Tuple[Any, ...]]) -> 'ArraySubClass': result = super().__getitem__(key) return result
重写后,pyright就能通过静态类型注解推断出切片返回的是ArraySubClass,不会再出现类型不匹配的提示。
补充说明
如果子类涉及布尔索引、整数数组索引等复杂场景,可以调整__getitem__的参数类型注解以贴合实际需求。此外,numpy的__array_function__和__array_ufunc__也可能存在类似的静态类型推断问题,必要时同样需要针对性添加类型注解或重写方法。
内容的提问来源于stack exchange,提问作者Kevlar
相关产品推荐
相关产品推荐

