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

如何基于Scikit-learn正确实现机器学习模型堆叠

模型堆叠构建原理与代码修正

一、模型堆叠核心原理

模型堆叠(Stacking)本质是分层建模:

  • 第一层(基模型):用多个不同模型对训练数据训练,每个基模型通过交叉验证生成训练集的预测结果(避免过拟合),这些预测结果会作为新特征。
  • 第二层(元模型):将基模型生成的新特征与原始特征(或仅用新特征)结合,训练一个顶层模型,最终由这个元模型输出最终预测结果。

你困惑的“合并各模型输出作为新特征”步骤,StackingRegressor已经自动完成:它会让每个基模型在训练集子集上拟合,再预测剩余子集,最终拼接成完整的训练集预测特征,喂给元模型训练。

二、现有代码的问题

  1. 数据不一致:每次实例化ModelWrapper都调用data_perp()生成新数据集,导致基模型和堆叠模型使用的数据不匹配。
  2. 数据泄漏:你实例化的rfw和xgbr已提前在训练集上拟合,而StackingRegressor需要未拟合的模型实例,它会自行处理基模型的训练流程,提前拟合会导致泄漏。
  3. 逻辑冗余:predict()方法里重复拟合模型,无意义。

三、修正后的代码

from sklearn.ensemble import RandomForestRegressor, StackingRegressor
from xgboost import XGBRegressor
from sklearn.metrics import mean_absolute_error, r2_score

# 统一生成数据,避免多次调用导致不一致
def data_perp():
    # 替换为你实际的数据处理逻辑
    from sklearn.datasets import fetch_california_housing
    from sklearn.model_selection import train_test_split
    from sklearn.preprocessing import StandardScaler
    
    data = fetch_california_housing()
    X = data.data
    y = data.target
    x_train, x_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
    
    scaler = StandardScaler()
    x_train_scaled = scaler.fit_transform(x_train)
    x_test_scaled = scaler.transform(x_test)
    
    return x_train, x_test, y_train, y_test, x_train_scaled, x_test_scaled

# 全局复用数据集
x_train, x_test, y_train, y_test, x_train_scaled, x_test_scaled = data_perp()

class ModelWrapper(object):
    def __init__(self, clf, params=None, stack=False, model_name=None):
        self.model_name = model_name
        self.params = params or {}
        # 仅初始化模型,不提前拟合
        self.clf = clf(**self.params)
        
        # 绑定全局数据集
        self.x_train = x_train
        self.x_test = x_test
        self.y_train = y_train
        self.y_test = y_test
        self.x_train_scaled = x_train_scaled
        self.x_test_scaled = x_test_scaled
        
        self.y_pred = None
        self.mae = None
        self.r2 = None
        
        # 非堆叠模式自动训练评估
        if not stack:
            self.train()
            self.predict()
            self.evaluate_model()

    def train(self):
        # 根据模型特性选择用原始数据或缩放数据训练
        self.clf.fit(self.x_train_scaled, self.y_train)

    def predict(self):
        self.y_pred = self.clf.predict(self.x_test_scaled)
        
    def evaluate_model(self):
        self.mae = mean_absolute_error(self.y_test, self.y_pred)
        self.r2 = r2_score(self.y_test, self.y_pred)
        print(f"-----------\n{self.model_name}")
        print(f"r2: {self.r2:.4f}")
        print(f"mae: {self.mae:.4f}\n-----------\n")

    def stack_predict(self, base_models):
        stacked_model = StackingRegressor(
            estimators=[(model.model_name, model.clf) for model in base_models],
            final_estimator=self.clf,
            cv=5,  # 交叉验证折数,控制基模型预测的稳定性
            passthrough=True  # 元模型输入包含原始特征+基模型预测特征
        )
        
        # 堆叠模型自动完成基模型交叉验证预测、元模型训练流程
        stacked_model.fit(self.x_train_scaled, self.y_train)
        self.y_pred = stacked_model.predict(self.x_test_scaled)
        self.evaluate_model()


# 实例化基模型:stack=True表示仅初始化不自动训练
forest_params = {'random_state':42}
xgboost_params = {'random_state':42, 'n_estimators':100}

rfw = ModelWrapper(clf=RandomForestRegressor, params=forest_params, stack=True, model_name='RandomForestRegressor')
xgbr = ModelWrapper(clf=XGBRegressor, params=xgboost_params, stack=True, model_name="XGBRegressor")

# 构建并训练堆叠模型
base_models = [rfw, xgbr]
stacked_model = ModelWrapper(clf=RandomForestRegressor, params=forest_params, stack=True, model_name='Stacked Model')
stacked_model.stack_predict(base_models)

四、关键说明

  • 数据统一:提前生成并复用数据集,确保所有模型使用同一批数据。
  • 基模型不提前拟合:通过stack=True让基模型仅初始化不训练,交给StackingRegressor处理训练,避免数据泄漏。
  • StackingRegressor参数:
    • cv:交叉验证折数,折数越高结果越稳定,但训练时间越长。
    • passthrough=True:元模型输入包含原始特征和基模型预测特征,设为False则仅用基模型预测特征。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 19:02:44