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

为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.0
  • torch==2.0.0
  • numpy==1.24.3
  • pandas==2.0.1
  • pandas-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的类型判断:

  1. 尝试升级pandas-stubs到最新版本,确保类型存根的准确性:

    pip install --upgrade pandas-stubs
    
  2. 显式使用类型的完整路径,避免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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 00:30:06