Scikit-learn Pipeline:构建可复用的 ML 流水线

Scikit-learn Pipeline:构建可复用的 ML 流水线

1. Pipeline 基础

python 复制代码
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
from sklearn.decomposition import PCA
from sklearn.ensemble import RandomForestClassifier

# 创建流水线
pipe = Pipeline([
    ('scaler', StandardScaler()),
    ('pca', PCA(n_components=10)),
    ('clf', RandomForestClassifier(n_estimators=100))
])

# 训练
pipe.fit(X_train, y_train)

# 预测
y_pred = pipe.predict(X_test)

# 评分
score = pipe.score(X_test, y_test)

2. ColumnTransformer

python 复制代码
from sklearn.compose import ColumnTransformer
from sklearn.preprocessing import StandardScaler, OneHotEncoder
from sklearn.impute import SimpleImputer

numeric_features = ['age', 'income', 'score']
categorical_features = ['city', 'gender']

preprocessor = ColumnTransformer([
    ('num', Pipeline([
        ('imputer', SimpleImputer(strategy='median')),
        ('scaler', StandardScaler()),
    ]), numeric_features),
    ('cat', Pipeline([
        ('imputer', SimpleImputer(strategy='most_frequent')),
        ('encoder', OneHotEncoder(handle_unknown='ignore')),
    ]), categorical_features),
])

# 完整流水线
pipe = Pipeline([
    ('preprocessor', preprocessor),
    ('classifier', RandomForestClassifier())
])

3. 自定义 Transformer

python 复制代码
from sklearn.base import BaseEstimator, TransformerMixin

class FeatureEngineer(BaseEstimator, TransformerMixin):
    def __init__(self, add_interaction=True):
        self.add_interaction = add_interaction
    
    def fit(self, X, y=None):
        return self
    
    def transform(self, X):
        X = X.copy()
        X['price_per_sqft'] = X['price'] / X['area']
        if self.add_interaction:
            X['age_income'] = X['age'] * X['income']
        return X

# 使用
pipe = Pipeline([
    ('feature_eng', FeatureEngineer()),
    ('scaler', StandardScaler()),
    ('clf', RandomForestClassifier())
])

4. GridSearch + Pipeline

python 复制代码
from sklearn.model_selection import GridSearchCV

param_grid = {
    'pca__n_components': [5, 10, 15],
    'clf__n_estimators': [50, 100, 200],
    'clf__max_depth': [5, 10, None],
}

grid = GridSearchCV(pipe, param_grid, cv=5, scoring='accuracy')
grid.fit(X_train, y_train)
print(f"最佳参数: {grid.best_params_}")

总结

组件 作用
Pipeline 串联处理步骤
ColumnTransformer 按列分别处理
FeatureUnion 并行特征提取
自定义 Transformer 封装业务逻辑
相关推荐
蒸鱼Yuzheng1 小时前
设备端性能工件可靠导出:断点续传、哈希、manifest 与失败恢复
android·自动化测试·python·adb·数据完整性
明志数科1 小时前
具身智能数据供给的分层:分布式采集与入厂采集的工程边界分析
人工智能·机器学习·机器人
圆圆讲门店2 小时前
挑选同城获客服务机构时需要考量的核心因素都有哪些?
大数据·网络·人工智能·python
傻啦嘿哟2 小时前
Python的默认参数把我坑惨了,原来写[]和写None的区别这么大
开发语言·python·机器学习
微小冷2 小时前
Python凸优化cvxpy初步
python·机器人·凸优化·数学规划·cvxpy
在世修行3 小时前
干货:显式映射 vs 自动判据
python·插件
xzal123 小时前
Python之简单理解栈和队列
python
lupai3 小时前
维修保养记录精准版 API 对接实战指南
数据库·python
罗西的思考3 小时前
[Agent Memory / 强化学习] MemPO源码学习笔记 ---(5)--- GRPO
人工智能·算法·机器学习
水水不水啊3 小时前
告别杂乱的调试窗口:我用 Python + WebView 写了一个现代化串口助手
python·测试工具·嵌入式·嵌入式开发·串口调试·串口助手