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

sklearn自定义Transformer加类型提示遇继承类类型错误

问题原因

你遇到的类型检查错误,是因为类型检查工具(如mypy)无法识别BaseEstimator和TransformerMixin的类型信息——要么是你的scikit-learn版本过低(1.2.0版本才开始为核心基类添加官方类型注解),要么是缺少对应的类型存根文件。另外,Self类型在Python 3.11及以上才归入标准typing模块,低于该版本需要单独导入。

解决方法

方法1:升级scikit-learn到最新稳定版

运行以下命令升级,让工具能识别官方提供的基类类型注解:

pip install --upgrade scikit-learn

方法2:安装第三方类型存根包(不升级版本时用)

如果无法升级scikit-learn,可安装第三方维护的类型存根补充类型信息:

pip install types-scikit-learn

方法3:兼容低版本Python的Self类型

若你的Python版本低于3.11,需从typing_extensions导入Self:

from typing_extensions import Self
from sklearn.base import BaseEstimator, TransformerMixin
import pandas as pd

class RemoveDuplicateRows(BaseEstimator, TransformerMixin):
    """Custom transformer for remove duplicate rows."""

    def fit(self, X: pd.DataFrame, y: pd.Series = None) -> Self:
        """Learn the parameters."""
        return self

    def transform(self, X: pd.DataFrame, y: pd.Series = None) -> pd.DataFrame:
        """Transform the input."""
        return X.drop_duplicates()

临时方案:添加类型忽略注释

如果以上方法都不想采用,可直接跳过类型检查工具的报错:

from sklearn.base import BaseEstimator, TransformerMixin
import pandas as pd

class RemoveDuplicateRows(BaseEstimator, TransformerMixin):  # type: ignore[misc]
    """Custom transformer for remove duplicate rows."""

    def fit(self, X: pd.DataFrame, y: pd.Series = None) -> "RemoveDuplicateRows":
        """Learn the parameters."""
        return self

    def transform(self, X: pd.DataFrame, y: pd.Series = None) -> pd.DataFrame:
        """Transform the input."""
        return X.drop_duplicates()

这里把Self替换为类名字符串,同时用注释跳过父类类型的错误检查。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 07:17:26