为Pandas、Torch、Numpy输入设置类型重载时解决Mypy报错问题
自定义模块重载类型定义的mypy报错解决
问题场景
我尝试为自定义模块的__call__方法设置重载类型,实现传入对应数据类型时返回相同类型,代码如下:
class MyModule: @overload def __call__(self, inputs: pd.DataFrame) -> pd.DataFrame: ... @overload def __call__(self, inputs: torch.Tensor) -> torch.Tensor: ... @overload def __call__(self, inputs: np.ndarray) -> np.ndarray: ... def __call__(self, inputs: np.ndarray | torch.Tensor | pd.DataFrame): pass
运行mypy时抛出以下错误:
utils.py:71: error: Overloaded function signature 2 will never be matched: signature 1's parameter type(s) are the same or broader [misc] utils.py:71: error: Overloaded function signature 3 will never be matched: signature 1's parameter type(s) are the same or broader [misc] utils.py:71: error: Overloaded function signature 3 will never be matched: signature 2's parameter type(s) are the same or broader [misc] Found 3 errors in 1 file (checked 21 source files)
这三个类型的实例在isinstance测试中不会互相返回True,我无法理解报错原因。
补充信息
使用的包版本:
mypy==1.2.0torch==2.0.0numpy==1.24.3pandas==2.0.1pandas-stubs==2.0.1.230501
我在文件中添加reveal_type(pd.DataFrame)、reveal_type(np.ndarray)和reveal_type(torch.Tensor)后运行mypy,得到以下输出(已翻译为中文):
utils.py:13: note: 推导类型为"Overload( def (data: Union[Union[typing.Sequence[Any], numpy.ndarray[Any, Any], pandas.core.series.Series[Any], pandas.core.indexes.base.Index], pandas.core.frame.DataFrame, builtins.dict[Any, Any], typing.Iterable[Union[Union[typing.Sequence[Any], numpy.ndarray[Any, Any], pandas.core.series.Series[Any], pandas.core.indexes.base.Index], Tuple[typing.Hashable, Union[typing.Sequence[Any], numpy.ndarray[Any, Any], pandas.core.series.Series[Any], pandas.core.indexes.base.Index]], builtins.dict[Any, Any]]], None] =, index: Union[Union[pandas.core.indexes.base.Index, pandas.core.series.Series[Any], numpy.ndarray[Any, Any], builtins.list[Any], builtins.dict[Any, Any], builtins.range, builtins.tuple[Any, ...]], None] =, columns: Union[Union[pandas.core.indexes.base.Index, pandas.core.series.Series[Any], numpy.ndarray[Any, Any], builtins.list[Any], builtins.dict[Any, Any], builtins.range, builtins.tuple[Any, ...]], None] =, dtype: Any =, copy: builtins.bool =) -> pandas.core.frame.DataFrame, def (data: Union[builtins.str, builtins.bytes, datetime.date, datetime.datetime, datetime.timedelta, numpy.datetime64, numpy.timedelta64, builtins.bool, builtins.int, builtins.float, pandas._libs.tslibs.timestamps.Timestamp, pandas._libs.tslibs.timedeltas.Timedelta, builtins.complex], index: Union[pandas.core.indexes.base.Index, pandas.core.series.Series[Any], numpy.ndarray[Any, Any], builtins.list[Any], builtins.dict[Any, Any], builtins.range, builtins.tuple[Any, ...]], columns: Union[pandas.core.indexes.base.Index, pandas.core.series.Series[Any], numpy.ndarray[Any, Any], builtins.list[Any], builtins.dict[Any, Any], builtins.range, builtins.tuple[Any, ...]], dtype: Any =, copy: builtins.bool =) -> pandas.core.frame.DataFrame)" utils.py:14: note: 推导类型为"def [_ShapeType <: Any, _DType_co <: numpy.dtype[Any]] (shape: Union[typing.SupportsIndex, typing.Sequence[typing.SupportsIndex]], dtype: Union[numpy.dtype[Any], None, Type[Any], numpy._typing._dtype_like._SupportsDType[numpy.dtype[Any]], builtins.str, Tuple[Any, builtins.int], Tuple[Any, Union[typing.SupportsIndex, typing.Sequence[typing.SupportsIndex]]], builtins.list[Any], TypedDict('numpy._typing._dtype_like._DTypeDict', {'names': typing.Sequence[builtins.str], 'formats': typing.Sequence[Any], 'offsets'?: typing.Sequence[builtins.int], 'titles'?: typing.Sequence[Any], 'itemsize'?: builtins.int, 'aligned'?: builtins.bool}), Tuple[Any, Any]] =, buffer: Union[builtins.bytes, builtins.bytearray, builtins.memoryview, array.array[Any], mmap.mmap, numpy.ndarray[Any, numpy.dtype[Any]], numpy.generic] =, offset: typing.SupportsIndex =, strides: Union[typing.SupportsIndex, typing.Sequence[typing.SupportsIndex]] =, order: Union[None, Literal['K'], Literal['A'], Literal['C'], Literal['F']] =) -> numpy.ndarray[_ShapeType`1, _DType_co`2]" utils.py:15: note: 推导类型为"Overload( def (*args: Any, *, device: Union[torch._C.device, builtins.str, builtins.int, None] =) -> torch._tensor.Tensor, def (storage: torch.types.Storage) -> torch._tensor.Tensor, def (other: torch._tensor.Tensor) -> torch._tensor.Tensor, def (size: Union[torch._C.Size, builtins.list[builtins.int], builtins.tuple[builtins.int, ...]], *, device: Union[torch._C.device, builtins.str, builtins.int, None] =) -> torch._tensor.Tensor)"
但错误仍然存在。
问题原因
mypy的报错源于类型存根定义的兼容性误判:尽管这三个类型在运行时不会互相兼容,但mypy解析类型存根时,认为其中某个类型的定义范围覆盖了其他类型(比如pd.DataFrame的类型签名被判定为可以兼容torch.Tensor或np.ndarray),导致后续的重载签名永远不会被匹配到。
解决方法
方案1:使用TypeVar绑定输入输出类型(推荐)
用TypeVar指定允许的类型范围,直接实现输入输出类型一致的约束,写法更简洁且避免重载冲突:
from typing import TypeVar, Union import pandas as pd import torch import numpy as np # 定义只能是指定三种类型的TypeVar T = TypeVar('T', pd.DataFrame, torch.Tensor, np.ndarray) class MyModule: def __call__(self, inputs: T) -> T: # 你的业务逻辑实现 pass
方案2:修正重载签名的类型标识
如果必须保留重载写法,可以通过显式指定类型的精确路径,或者升级pandas-stubs版本来修正mypy的类型判断:
尝试升级
pandas-stubs到最新版本,确保类型存根的准确性:pip install --upgrade pandas-stubs显式使用类型的完整路径,避免mypy误判:
from typing import overload import pandas as pd import torch import numpy as np class MyModule: @overload def __call__(self, inputs: pd.core.frame.DataFrame) -> pd.core.frame.DataFrame: ... @overload def __call__(self, inputs: torch._tensor.Tensor) -> torch._tensor.Tensor: ... @overload def __call__(self, inputs: np.ndarray) -> np.ndarray: ... def __call__(self, inputs: np.ndarray | torch._tensor.Tensor | pd.core.frame.DataFrame): pass
内容的提问来源于stack exchange,提问作者iHowell
相关产品推荐
相关产品推荐

