如何结合绑定类型变量与静态类型实现灵活的Python类型提示?
解决numpy类数组函数的类型提示冲突问题
问题背景
要为一个仅依赖numpy实现的函数添加类型提示,要求:
- 输入支持所有
numpy.typing.ArrayLike类型的类数组对象 - 返回规则:输入为列表等普通类数组时,返回
np.NDArray;输入为pandas.DataFrame/Series时,返回原输入类型
现有方案的问题
方案1:固定返回np.NDArray
import numpy as np import numpy.typing as npt import pandas as pd def get_decibels1(p2: npt.ArrayLike) -> npt.NDArray: return 10 * np.log10(np.divide(p2, 4e-10)) df = pd.DataFrame([[4, 5, 6], [7, 8, 9]]) get_decibels1(df).columns # mypy报错:"ndarray[Any, dtype[Any]]" 没有属性"columns"
问题:无法识别pandas输入的返回类型,导致调用pandas专属方法时触发类型检查错误。
方案2:绑定ArrayLike的TypeVar
import numpy as np import numpy.typing as npt import pandas as pd from typing import TypeVar T = TypeVar('T', bound=npt.ArrayLike) def get_decibels2(p2: T) -> T: return 10 * np.log10(np.divide(p2, 4e-10)) ls = [4.0, 5, 6] get_decibels2(ls).shape # mypy报错:"list[float]" 没有属性"shape"
问题:numpy会将列表等普通类数组转换为NDArray,但类型提示错误地认为返回原列表类型,触发属性不存在的错误。
方案3:直接重载(签名重叠)
import numpy as np import numpy.typing as npt import pandas as pd from typing import TypeVar, overload, Union T = TypeVar('T', bound=Union[pd.DataFrame, pd.Series]) @overload def get_decibels(p2: T) -> T: ... @overload def get_decibels(p2: npt.ArrayLike) -> npt.NDArray: ... def get_decibels(p2: npt.ArrayLike): return 10 * np.log10(np.divide(p2, 4e-10))
问题:mypy报错Overloaded function signatures 1 and 2 overlap with incompatible return types,因为pd.DataFrame属于npt.ArrayLike,两个签名的输入范围重叠,无法区分。
正确的重载实现
通过排除法明确区分两类输入的重载签名,避免范围重叠:
- 第一个重载明确匹配pandas的DataFrame和Series
- 第二个重载匹配**非pandas的
ArrayLike**类型
import numpy as np import numpy.typing as npt import pandas as pd from typing import TypeVar, overload, Exclude # 定义绑定pandas类型的TypeVar T = TypeVar('T', bound=pd.DataFrame | pd.Series) @overload def get_decibels(p2: T) -> T: ... @overload def get_decibels(p2: Exclude[npt.ArrayLike, pd.DataFrame | pd.Series]) -> npt.NDArray: ... def get_decibels(p2: npt.ArrayLike): return 10 * np.log10(np.divide(p2, 4e-10))
验证示例
# 测试pandas输入 df = pd.DataFrame([[4, 5, 6], [7, 8, 9]]) print(get_decibels(df).columns) # mypy无报错,正确识别返回DataFrame # 测试列表输入 ls = [4.0, 5, 6] print(get_decibels(ls).shape) # mypy无报错,正确识别返回NDArray # 测试numpy数组输入 arr = np.array([4.0, 5, 6]) print(get_decibels(arr).shape) # mypy无报错,正确识别返回NDArray
原理说明
Exclude[npt.ArrayLike, pd.DataFrame | pd.Series]精确限定了第二个重载的输入范围,与第一个重载的pandas类型完全不重叠- mypy会根据输入的具体类型匹配对应的重载签名,返回正确的类型提示
内容的提问来源于stack exchange,提问作者ti-sch
相关产品推荐
相关产品推荐

