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

如何在tsai框架fit()训练过程中实现Early Stopping早停

问题根源

你选错回调函数了。TerminateOnNaNCallback() 没有早停逻辑,它只会在训练过程中出现损失值为NaN的数值溢出异常时强制终止训练,完全不会监控验证集指标变化,自然实现不了你要的早停效果。

tsai基于fastai开发,直接用fastai内置的EarlyStoppingCallback就能实现标准早停,不需要自己写额外逻辑。

正确配置方法

早停回调有三个核心参数需要根据任务调整:

  • monitor:指定要监控的指标名称,注意fastai中验证集指标默认带valid_前缀,比如你代码里用了mae、rmse两个指标,监控验证集MAE就填'valid_mae',监控验证集RMSE就填'valid_rmse'
  • patience:容忍轮次,即连续多少个epoch监控指标没有提升就停止训练,一般设10-20比较常用
  • min_delta:判定指标提升的最小阈值,波动幅度小于这个值不算有效提升,避免训练过程中微小的指标震荡触发不必要的停止

注意你初始化TSRegressor时已经传入了ShowGraph()回调,不要在fit阶段单独传单个回调把原有回调覆盖,要么初始化learner时就把早停加进回调列表,要么fit时把需要的所有回调都放在列表里传入。

修改后的参考代码

推荐在初始化learner时就配置好所有回调,避免遗漏:

from tsai.all import *

dsid = 'AppliancesEnergy'
arch_config = {
    'hidden_size':100, 
    'n_layers':2, 
    'rnn_dropout':0.2, 
    'fc_dropout':0.5, 
    'bidirectional':True
}

X, y, splits = get_regression_data(dsid, split_data=False)
learn = TSRegressor(
    X, 
    y, 
    splits=splits, 
    bs=128, 
    batch_tfms=[TSStandardize(by_sample=True)], 
    arch=LSTM, 
    arch_config=arch_config, 
    metrics=[mae, rmse], 
    cbs=[
        ShowGraph(),
        # 早停配置:监控验证集MAE,连续15轮无提升就停止,最小提升阈值0.001
        EarlyStoppingCallback(monitor='valid_mae', patience=15, min_delta=0.001)
    ], 
    verbose=True
)

learn.fit_one_cycle(100, lr_max=1e-3)
learn.plot_metrics()

小提示:mae、rmse都是越小越好的指标,EarlyStoppingCallback会自动识别优化方向,不需要额外传参。如果是监控准确率这类越大越好的指标,才需要额外传入comp=np.greater指定比较规则。触发早停后,框架会自动加载监控指标表现最好的那一轮的模型权重,不需要手动保存和恢复。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 05:45:37