如何为调用fit方法后生成的属性使用Protocol协议?
问题描述
我正在开发一个深度使用scikit-learn对象的包,希望借助Protocol定义sklearn部分功能的接口。例如,对象调用fit方法后会生成feature_names_in_属性,但该属性在执行fit前不存在,当前的Protocol定义会因为实例化时属性不存在而报错,代码示例如下:
_TransformerSelf = TypeVar('_TransformerSelf', bound='TransformerMixin') @runtime_checkable class P(Protocol): # 实例化时该属性不存在,导致报错 feature_names_in_: Sequence[str] def fit( self: _TransformerSelf, X: pd.DataFrame, y: pd.DataFrame | None = None ) -> _TransformerSelf: ...
解决方案
针对这种"拟合后才生成属性"的场景,可以通过以下两种方式调整Protocol的定义,适配scikit-learn的对象生命周期:
方案一:可选属性+类型守卫
将feature_names_in_标记为可选属性,同时编写一个类型守卫函数,用于判断对象是否已经完成拟合。这样类型检查器可以在调用fit后识别出属性已存在:
from typing import Protocol, TypeVar, Sequence, Optional, TypeGuard import pandas as pd from sklearn.base import TransformerMixin _TransformerSelf = TypeVar('_TransformerSelf', bound='TransformerMixin') class TransformerProtocol(Protocol): feature_names_in_: Optional[Sequence[str]] # 标记为可选属性 def fit( self: _TransformerSelf, X: pd.DataFrame, y: Optional[pd.DataFrame] = None ) -> _TransformerSelf: ... def is_fitted(transformer: TransformerProtocol) -> TypeGuard[TransformerProtocol & {"feature_names_in_": Sequence[str]}]: return hasattr(transformer, 'feature_names_in_') and transformer.feature_names_in_ is not None # 使用示例 def process_fitted_transformer(transformer: TransformerProtocol): if is_fitted(transformer): # 此处类型检查器会确认feature_names_in_为Sequence[str]类型 print(transformer.feature_names_in_) else: raise ValueError("Transformer尚未拟合,请先调用fit方法")
方案二:拆分拟合前后的Protocol
定义两个Protocol,分别对应未拟合和已拟合的状态,让fit方法返回已拟合的类型,明确区分对象的不同生命周期阶段:
from typing import Protocol, TypeVar, Sequence import pandas as pd from sklearn.base import TransformerMixin _TransformerSelf = TypeVar('_TransformerSelf', bound='TransformerMixin') class UnfittedTransformerProtocol(Protocol): def fit( self: _TransformerSelf, X: pd.DataFrame, y: Optional[pd.DataFrame] = None ) -> FittedTransformerProtocol: ... class FittedTransformerProtocol(UnfittedTransformerProtocol, Protocol): feature_names_in_: Sequence[str] # 使用示例 def fit_transformer(transformer: UnfittedTransformerProtocol, X: pd.DataFrame) -> FittedTransformerProtocol: return transformer.fit(X) fitted_transformer = fit_transformer(SomeSklearnTransformer(), X_data) # 此处类型检查器会自动识别feature_names_in_已存在 print(fitted_transformer.feature_names_in_)
注意事项
- 如果使用
@runtime_checkable装饰器,要确保运行时检查逻辑符合预期:比如方案一中的类型守卫可以配合运行时检查,避免未拟合对象进入需要已拟合属性的逻辑。 - 第二种方案更贴合scikit-learn对象的状态区分,类型语义更清晰,适合严格的类型检查场景。
内容的提问来源于stack exchange,提问作者Collin Cunningham
相关产品推荐
相关产品推荐

