调用np.argmin返回值触发mypy类型错误,如何正确注解函数
问题原因
np.argmin的静态类型注解为重载定义:未指定axis或keepdims参数时,返回值为numpy整数标量,只有当指定axis或设置keepdims=True时才会返回np.ndarray类型。你当前的调用方式返回的是标量,和标注的np.ndarray不匹配,因此mypy抛出错误。
正确注解方案
方案1:确认返回为标量时直接标注对应类型
如果你确定调用np.argmin时不需要返回数组,直接将返回值标注为整数类型即可,numpy标量整数完全兼容原生int类型:
import numpy as np def test() -> int: return np.argmin(np.array([1, 2, 3]))
如果需要严格标注为numpy原生标量类型,可以导入numpy.typing下的对应类型:
import numpy as np from numpy.typing import Intp def test() -> Intp: return np.argmin(np.array([1, 2, 3]))
方案2:需要返回数组时补充keepdims参数
如果你确实需要返回np.ndarray类型,给np.argmin添加keepdims=True参数,此时返回值为维度保留的数组,匹配你最初的类型标注:
import numpy as np from numpy.typing import NDArray def test() -> NDArray[np.intp]: return np.argmin(np.array([1, 2, 3]), keepdims=True)
方案3:兼容两种返回场景的重载注解
如果你的函数会根据入参决定返回值类型,可以用typing.overload做重载注解覆盖两种情况:
import numpy as np from typing import overload, Optional, Literal from numpy.typing import NDArray, Intp @overload def test(keepdims: Literal[True]) -> NDArray[np.intp]: ... @overload def test(keepdims: Literal[False] = ...) -> Intp: ... def test(keepdims: bool = False) -> Intp | NDArray[np.intp]: return np.argmin(np.array([1, 2, 3]), keepdims=keepdims)
内容的提问来源于stack exchange,提问作者dmmpie
相关产品推荐
相关产品推荐

