为numpy ndarray编写类型提示报_DTypeMeta不可下标如何解决
运行时抛出该错误的核心原因是:Python 3.10 以前版本的类型注解默认会在运行时求值,而 numpy 运行时的dtype、ndarray类型本身不支持泛型下标语法,所以会触发类型不可下标的报错。
首先在代码顶部加入如下语句关闭注解的运行时求值,即可避免该报错:
from __future__ import annotations
如果使用Python 3.6及更早版本无法使用上述语法,可将所有ndarray相关的标注用英文引号包裹,同样可以避免运行时求值报错。
你可以根据需求选择以下两种标注方案:
方案1:numpy内置标注(无第三方依赖)
numpy 1.21及以上版本内置了numpy.typing模块,可直接用于标注数组的dtype:
from __future__ import annotations import numpy as np from numpy.typing import NDArray # 定义元素类型为uint8的数组类型 RGBArray = NDArray[np.uint8] def load_images(paths: list[str]) -> tuple[list[RGBArray], list[str]]: # 业务逻辑实现 ...
如果需要标注维度信息,可结合typing.Annotated自定义维度标记:
from typing import Annotated # 标记为3维uint8类型的数组 ThreeDRGBArray = Annotated[NDArray[np.uint8], "shape=(H, W, 3)"]
方案2:nptyping精准标注(支持形状校验)
如果需要静态检查器识别数组的形状规则,可安装第三方库nptyping实现更灵活的标注:
from __future__ import annotations import numpy as np from nptyping import NDArray, Shape, UInt8 def load_images(paths: list[str]) -> tuple[list[NDArray[Shape["*, *, 3"], UInt8]], list[str]]: # 标注含义:任意高度、任意宽度、通道数为3的uint8类型数组 ...
内容的提问来源于stack exchange,提问作者palapapa
相关产品推荐
相关产品推荐

