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
相关产品推荐
相关产品推荐

