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

如何为调用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 08:40:14