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

Keras LSTM训练早停:restore_best_weights参数未按预期工作问题

关于Keras EarlyStopping中restore_best_weights未按预期工作的问题

我尝试训练一个LSTM网络,用Keras的callbacks模块实现早停,代码如下:

callback = tensorflow.keras.callbacks.EarlyStopping(monitor='loss', min_delta=0.0001, patience=7, mode='min', restore_best_weights=True, verbose=1)
model1= Sequential()
model1.add(LSTM(64, activation='swish',input_shape=(trainX.shape[1], trainX.shape[2]),return_sequences=True))
model1.add(LSTM(128,activation = 'swish', return_sequences=True))
model1.add(LSTM(64,activation = 'elu', return_sequences=False))
model1.add(Dropout(0.01))
model1.add(Dense(trainY.shape[1]))
model1.compile(optimizer='adam', loss='mse')
model1.summary()
model1.fit(trainX,trainY, epochs=n_epochs, batch_size=batchsize, verbose=2, callbacks=[callback])

但发现restore_best_weights参数没按预期工作:设置为True后,早停触发却没加载loss最低轮次的权重。训练日志如下:

Epoch 1/9
1250/1250 - 76s - loss: 0.0012 - 76s/epoch - 61ms/step
Epoch 2/9
1250/1250 - 76s - loss: 0.0011 - 76s/epoch - 61ms/step
Epoch 3/9
1250/1250 - 76s - loss: 0.0011 - 76s/epoch - 60ms/step
Epoch 4/9
1250/1250 - 76s - loss: 0.0010 - 76s/epoch - 60ms/step
Epoch 5/9
1250/1250 - 76s - loss: 9.9930e-04 - 76s/epoch - 61ms/step
Epoch 6/9
1250/1250 - 75s - loss: 9.9933e-04 - 75s/epoch - 60ms/step
Epoch 7/9
Restoring model weights from the end of the best epoch: 3.
1250/1250 - 76s - loss: 0.0010 - 76s/epoch - 61ms/step
Epoch 7: early stopping

我预期加载第5轮的权重(loss最低),但实际恢复的是第3轮的权重,后续训练也没明显提升。请问是操作有误还是对参数理解错了?有没有更好的方法确保早停时选择loss最优的权重?


问题原因分析

1. min_delta参数设置过严

你设置的min_delta=0.0001,Keras的逻辑是:只有当监控指标(这里是训练loss)的下降幅度严格大于该值时,才会被判定为有效改进。从训练日志看:

  • 第3轮loss为0.0011,第4轮为0.0010,差值刚好等于0.0001,不满足“严格大于”的要求,因此不会更新最佳epoch;
  • 第5轮loss为0.000993,和第3轮的差值为0.000107,理论上满足条件,但可能因为训练loss的浮点精度波动,导致模型未识别到该有效改进,最终最佳epoch停留在第3轮。

2. 仅监控训练集loss的局限性

你监控的是训练集loss,而训练过程中训练loss可能出现小幅波动,甚至会因为过拟合出现下降后回升的情况,无法准确反映模型的泛化能力,也容易导致最佳epoch判断偏差。


解决办法

1. 调整min_delta阈值

缩小min_delta的值,让小幅的loss下降也能被识别为有效改进:

callback = tensorflow.keras.callbacks.EarlyStopping(
    monitor='loss', 
    min_delta=1e-5,  # 降低阈值,适配小幅度loss变化
    patience=7, 
    mode='min', 
    restore_best_weights=True, 
    verbose=1
)

2. 改用验证集loss监控(推荐)

划分验证集,监控val_loss,能更合理判断模型的泛化能力,避免过拟合带来的判断偏差:

# 调整EarlyStopping监控指标为验证集loss
callback = tensorflow.keras.callbacks.EarlyStopping(
    monitor='val_loss', 
    min_delta=1e-5, 
    patience=7, 
    mode='min', 
    restore_best_weights=True, 
    verbose=1
)

# 在fit中加入验证集,比如用10%训练数据做验证
model1.fit(
    trainX,trainY, 
    epochs=n_epochs, 
    batch_size=batchsize, 
    verbose=2, 
    callbacks=[callback],
    validation_split=0.1
)

3. 结合ModelCheckpoint双重保障

搭配ModelCheckpoint回调手动保存最佳权重,作为restore_best_weights的补充,确保权重不会丢失:

from tensorflow.keras.callbacks import ModelCheckpoint

# 保存最佳权重到本地文件
checkpoint = ModelCheckpoint(
    'best_lstm_weights.h5',
    monitor='val_loss',
    save_best_only=True,
    mode='min',
    verbose=1
)

# 同时使用两个回调
model1.fit(
    trainX,trainY, 
    epochs=n_epochs, 
    batch_size=batchsize, 
    verbose=2, 
    callbacks=[callback, checkpoint],
    validation_split=0.1
)

# 训练结束后可手动加载最佳权重
model1.load_weights('best_lstm_weights.h5')

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 05:30:49