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

如何为支持数值或类数组输入的NumPy操作函数正确添加类型提示?

Python类型提示:处理数值与NumPy数组兼容的最佳实践

问题背景

当函数需要同时接受普通数值(int/float)和NumPy数组作为参数时,直接标注np.ndarray会导致类型检查工具(如mypy)报错,因为普通数值类型不匹配ndarray类型。例如:

import numpy as np

def square(x: np.ndarray):
    return x**2

num = 9
print(square(num))  # mypy报错:int类型与ndarray不兼容
arr = np.array([3, 6.31, 9, 8.73])
print(square(arr))

mypy错误提示:

error: Argument 1 to "square" has incompatible type "int"; expected "ndarray[Any, Any]"  [arg-type]

你尝试的临时方案通过Union[float, np.ndarray[Any, Any]]让mypy通过,但存在不够严谨的问题(比如未显式支持int类型,且Any会丢失数组的维度和dtype信息)。

最佳实践方案

1. 使用精确的Union类型配合numpy官方类型别名

从numpy.typing导入NDArray(替代直接用np.ndarray),结合具体数值类型构建Union,同时显式支持int、float:

import numpy as np
from numpy.typing import NDArray
from typing import Union

NumOrArray = Union[int, float, NDArray[np.number]]

def square(x: NumOrArray) -> NumOrArray:
    return x**2

# 测试代码
num = 9
print(square(num))  # 类型检查通过
arr = np.array([3, 6.31, 9, 8.73])
print(square(arr))  # 类型检查通过

优势:

  • NDArray[np.number]涵盖所有数值类型的NumPy数组,比NDArray[Any, Any]更精确
  • 显式包含int、float,避免类型检查遗漏
  • 同时标注返回值类型,让类型信息更完整

2. 使用抽象数值类型简化Union

如果不需要区分int和float,可以用numbers.Real抽象类型来涵盖所有实数类型,配合NDArray:

import numpy as np
from numpy.typing import NDArray
from typing import Union
import numbers

NumOrArray = Union[numbers.Real, NDArray[np.number]]

def square(x: NumOrArray) -> NumOrArray:
    return x**2

优势:

  • 用抽象类型减少重复,numbers.Real包含int、float、Decimal等所有实数类型
  • 保持类型检查的严谨性

3. 针对数组维度/dtype的精确标注

如果需要更严格的类型约束(比如只接受一维浮点数组),可以指定NDArray的参数:

import numpy as np
from numpy.typing import NDArray
from typing import Union

NumOr1DFloatArray = Union[int, float, NDArray[np.float64]]

def square(x: NumOr1DFloatArray) -> NumOr1DFloatArray:
    return x**2

4. 类型守卫(针对需要分支处理的场景)

如果函数内部需要对数值和数组做不同逻辑处理,可以用类型守卫来明确区分类型:

import numpy as np
from numpy.typing import NDArray
from typing import Union, TypeGuard

def is_ndarray(x: Union[int, float, NDArray]) -> TypeGuard[NDArray]:
    return isinstance(x, np.ndarray)

def square(x: Union[int, float, NDArray]) -> Union[int, float, NDArray]:
    if is_ndarray(x):
        # 此处x会被类型检查器识别为NDArray
        return x**2 + np.ones_like(x)
    else:
        # 此处x会被识别为int/float
        return x**2

优势:

  • 让类型检查器能准确推断分支内的变量类型,避免类型错误
  • 适合复杂逻辑的函数

总结

你的临时方案思路是对的,但可以通过以下方式优化:

  • 用numpy.typing.NDArray替代np.ndarray[Any, Any],获得更精确的数组类型信息
  • 显式包含int类型(或用numbers.Real简化)
  • 补充返回值的类型提示
  • 复杂场景下使用类型守卫

内容的提问来源于stack exchange,提问作者crabulus_maximus

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 18:27:53