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

如何在Keras+Scikit-Learn GridSearchCV中传递KFold验证集至model.fit

解决方案

要让KFold的验证集传入Keras模型的fit(validation_data)参数,从而支持基于val_loss的早停和学习曲线绘制,核心是绕开GridSearchCV默认不传递验证集给模型训练流程的限制,以下是两种可行方案:

方案一:手动实现网格搜索+交叉验证(推荐)

直接遍历参数组合和KFold拆分,在每个fold训练时显式传递验证集给model.fit,完全控制训练流程:

import pandas as pd
import tensorflow as tf
from scikeras.wrappers import KerasClassifier
from sklearn.model_selection import KFold, ParameterGrid
from sklearn.metrics import mean_squared_error, accuracy_score

# 1. 定义参数网格和KFold拆分
param_grid = NN_MONK1_GRID_DICT
kf = KFold(n_splits=5, shuffle=True, random_state=15)  # 替换为你的CV配置

# 2. 存储所有参数组合的交叉验证结果
results = []

# 3. 遍历所有参数组合
for params in ParameterGrid(param_grid):
    fold_metrics = {
        "train_loss": [], "val_loss": [],
        "accuracy": [], "mse": []
    }
    
    # 遍历每个KFold拆分
    for train_idx, val_idx in kf.split(X_train, y_train):
        X_tr, y_tr = X_train.iloc[train_idx], y_train.iloc[train_idx]
        X_val, y_val = X_train.iloc[val_idx], y_train.iloc[val_idx]
        
        # 初始化KerasClassifier,传入当前参数
        nn = KerasClassifier(
            model=get_NN,
            X_len=len(X_train.columns),
            loss="mse",
            optimizer="SGD",
            epochs=300,
            batch_size=4,
            shuffle=True,
            verbose=False,
            callbacks=[
                tf.keras.callbacks.EarlyStopping(
                    monitor="val_loss", min_delta=0.0001, patience=15, restore_best_weights=True
                )
            ]
        )
        
        # 为当前实例设置网格参数
        for param_key, param_val in params.items():
            setattr(nn, param_key, param_val)
        
        # 训练模型,显式传入验证集
        history = nn.fit(X_tr, y_tr, validation_data=(X_val, y_val))
        
        # 记录训练/验证损失
        fold_metrics["train_loss"].append(history.history["loss"][-1])
        fold_metrics["val_loss"].append(history.history["val_loss"][-1])
        
        # 计算并记录评估指标
        y_pred = nn.predict(X_val)
        fold_metrics["accuracy"].append(accuracy_score(y_val, y_pred.round()))  # 分类问题需调整预测值格式
        fold_metrics["mse"].append(mean_squared_error(y_val, y_pred))
    
    # 计算当前参数组合的平均指标
    avg_metrics = {k: sum(v)/len(v) for k, v in fold_metrics.items()}
    avg_metrics["params"] = params
    results.append(avg_metrics)

# 4. 找到最优参数(按mse最小筛选)
results_df = pd.DataFrame(results)
best_params = results_df.loc[results_df["mse"].idxmin(), "params"]

关键点说明:

  • 每个fold训练时直接将拆分好的验证集传入fit(validation_data),早停回调能正常监控val_loss
  • 可以直接从history对象提取训练/验证损失,用于绘制学习曲线
  • 完全自定义指标计算和结果存储,灵活性更高

方案二:自定义KerasClassifier子类结合PredefinedSplit

通过自定义estimator获取交叉验证的验证集,适配GridSearchCV的流程:

from scikeras.wrappers import KerasClassifier
from sklearn.model_selection import PredefinedSplit
import numpy as np

# 1. 自定义KerasClassifier子类,支持从全局数据集提取验证集
class KerasCVClassifier(KerasClassifier):
    def __init__(self, *args, X_full=None, y_full=None, **kwargs):
        super().__init__(*args, **kwargs)
        self.X_full = X_full
        self.y_full = y_full

    def fit(self, X, y, **kwargs):
        # 根据训练集索引找到对应的验证集索引
        train_indices = X.index
        val_indices = self.X_full.index.difference(train_indices)
        X_val = self.X_full.loc[val_indices]
        y_val = self.y_full.loc[val_indices]
        # 传递验证集给model.fit
        kwargs["validation_data"] = (X_val, y_val)
        return super().fit(X, y, **kwargs)

# 2. 生成PredefinedSplit所需的测试折叠标记
kf = KFold(n_splits=5, shuffle=True, random_state=15)
test_fold = np.full(len(X_train), -1)
for fold_idx, (_, val_idx) in enumerate(kf.split(X_train)):
    test_fold[val_idx] = fold_idx
ps = PredefinedSplit(test_fold)

# 3. 初始化自定义estimator
nn = KerasCVClassifier(
    model=get_NN,
    X_full=X_train,
    y_full=y_train,
    X_len=len(X_train.columns),
    loss="mse",
    optimizer="SGD",
    epochs=300,
    batch_size=4,
    shuffle=True,
    verbose=False,
    callbacks=[
        tf.keras.callbacks.EarlyStopping(
            monitor="val_loss", min_delta=0.0001, patience=15, restore_best_weights=True
        )
    ]
)

# 4. 运行GridSearchCV
grid = GridSearchCV(
    nn,
    param_grid=NN_MONK1_GRID_DICT,
    scoring=custom_scores_monk,
    refit="mse",
    cv=ps,
    return_train_score=True,
    n_jobs=-1
)
grid.fit(X_train, y_train)

关键点说明:

  • 自定义类通过全局数据集X_full和训练集索引,自动定位当前fold的验证集
  • 使用PredefinedSplit让GridSearchCV遵循我们定义的KFold拆分规则
  • 保留了GridSearchCV的便捷性,但需要注意数据集索引的一致性(建议X_train使用默认整数索引)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 15:35:56