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

如何在LSTM模型中结合Cross Validation与Early Stopping机制?

为LSTM模型实现基于3折交叉验证的全局早停机制

需求回顾

需要为LSTM模型结合3折交叉验证与早停机制:当连续3个epoch的3折平均AUC无提升时停止训练,最终构建具备最高3折平均AUC的模型。

修改后的完整代码

import numpy as np
from keras.models import Sequential
from keras.layers import LSTM, Dense
from keras.optimizers import Adam
from sklearn.model_selection import KFold
from sklearn.metrics import roc_auc_score

# 定义模型构建函数,方便重复初始化独立模型
def build_model(input_shape):
    model = Sequential()
    model.add(LSTM(units=32, input_shape=input_shape, dropout=0.3))
    model.add(Dense(1, activation='sigmoid'))
    model.compile(loss='binary_crossentropy', optimizer=Adam(learning_rate=0.01), metrics=['AUC'])
    return model

# 初始化核心参数
kfold = KFold(n_splits=3, shuffle=True)
input_shape = (X_train.shape[1], X_train.shape[2])
max_epochs = 100
patience = 3  # 连续3轮无提升触发早停
best_avg_auc = -np.inf
best_epoch = 0
best_model_weights = None
no_improve_count = 0

# 外层循环控制训练epoch,每轮遍历所有3折完成训练与评估
for epoch in range(max_epochs):
    fold_aucs = []
    # 遍历每个交叉验证折,训练并计算当前折的验证AUC
    for train_idx, val_idx in kfold.split(X_train):
        X_train_fold, X_val_fold = X_train[train_idx], X_train[val_idx]
        y_train_fold, y_val_fold = y_train[train_idx], y_train[val_idx]
        
        # 每折初始化新模型,保证交叉验证的独立性
        model = build_model(input_shape)
        # 训练至当前总epoch数(累计训练epoch+1轮)
        model.fit(X_train_fold, y_train_fold, epochs=epoch+1, batch_size=32, verbose=0)
        
        # 计算当前折的验证AUC
        y_pred = model.predict(X_val_fold, verbose=0)
        auc = roc_auc_score(y_val_fold, y_pred)
        fold_aucs.append(auc)
    
    # 计算当前epoch的3折平均AUC
    current_avg_auc = np.mean(fold_aucs)
    print(f"Epoch: {epoch+1}, 3折平均AUC: {current_avg_auc:.4f}")
    
    # 判断平均AUC是否提升,更新最优状态
    if current_avg_auc > best_avg_auc:
        best_avg_auc = current_avg_auc
        best_epoch = epoch + 1
        no_improve_count = 0
        # 用全量训练数据训练到最优epoch,保存权重
        best_model = build_model(input_shape)
        best_model.fit(X_train, y_train, epochs=best_epoch, batch_size=32, verbose=0)
        best_model_weights = best_model.get_weights()
    else:
        no_improve_count += 1
        # 触发早停条件,终止训练
        if no_improve_count >= patience:
            print(f"连续{patience}轮3折平均AUC无提升,提前停止训练")
            break

# 加载最优模型权重,得到最终模型
final_model = build_model(input_shape)
final_model.set_weights(best_model_weights)
print(f"最优模型训练完成,对应Epoch: {best_epoch},最优3折平均AUC: {best_avg_auc:.4f}")

关键改动说明

  • 模型构建标准化:封装build_model函数,确保每次交叉验证折都使用全新初始化的模型,避免不同折之间的参数污染。
  • 全局早停逻辑:将epoch循环置于外层,每轮遍历所有3折完成训练与评估,基于3折平均AUC判断是否提升,连续3轮无提升则停止,替代原代码中单折独立早停的逻辑。
  • 最优模型留存:每当平均AUC刷新记录时,用全量训练数据训练到当前epoch数并保存权重,最终模型直接加载最优权重,保证性能最优。
  • 依赖修复:补充原代码缺失的Adam优化器导入。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 21:19:52