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

如何在tsai中保存时间序列训练的最佳模型?

在tsai中保存训练周期内的最佳模型

tsai基于fastai框架,你可以直接用fastai的SaveModelCallback回调函数实现需求——自动保存验证损失最低的模型,和Keras的ModelCheckpoint功能完全对应。

修改后的完整代码:

import os
os.chdir(os.path.dirname(os.path.abspath(__file__)))
from pickle import load
from multiprocessing import Process
import numpy as np
from tsai.all import *
import matplotlib.pyplot as plt
from sklearn.metrics import precision_recall_curve

dataset_idx = 0
X_train = load(open(r"X_train_"+str(dataset_idx)+".pkl", 'rb'))
y_train = load(open(r"y_train_"+str(dataset_idx)+".pkl", 'rb'))
X_test = load(open(r"X_test_"+str(dataset_idx)+".pkl", 'rb'))
y_test = load(open(r"y_test_"+str(dataset_idx)+".pkl", 'rb'))
print("dataset loaded")

learn = TSClassifier(X_train, y_train, arch=InceptionTimePlus, arch_config=dict(fc_dropout=0.5))

print("training started")
# 添加SaveModelCallback回调,监控验证损失并保存最优模型
learn.fit_one_cycle(5, 0.0005, 
                    cbs=SaveModelCallback(monitor='valid_loss', mode='min', fname=f"best_tsai_{dataset_idx}"))

# 若需导出包含最优模型的完整learner,执行以下步骤
learn.load(f"best_tsai_{dataset_idx}")  # 先加载保存的最优权重
learn.export(f"tsai_best_{dataset_idx}.pkl")  # 导出完整learner

关键参数说明:

  • monitor='valid_loss':指定监控指标为验证损失
  • mode='min':因为要保留验证损失最小的模型,所以设为min;如果是监控准确率这类需要最大化的指标,改为max即可
  • fname:设置保存的模型权重文件名(默认存储在当前目录的models文件夹下)

后续加载最优模型:

训练结束后,可通过两种方式加载最优模型:

# 方式1:直接加载完整导出的learner
learn = load_learner(f"tsai_best_{dataset_idx}.pkl")

# 方式2:先初始化learner,再加载权重
learn = TSClassifier(X_train, y_train, arch=InceptionTimePlus, arch_config=dict(fc_dropout=0.5))
learn.load(f"best_tsai_{dataset_idx}")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 17:47:25