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

如何用MLflow保存加载Darts TFT模型?加载报错求助

解决MLflow保存/加载Darts TFT模型的问题

问题原因

你用sklearn flavor保存TFT模型是错误的,因为TFT是基于PyTorch Lightning的时序模型,不属于scikit-learn生态,sklearn的序列化逻辑无法正确处理TFT内部的Trainer等PyTorch Lightning组件,导致加载时出现AttributeError: 'Trainer' object has no attribute '_accelerator_connector'。

可行解决方案

方案1:使用MLflow PyTorch flavor保存加载

直接用mlflow.pytorch的API来处理TFT模型,因为TFT本质是PyTorch模型:

保存模型

import mlflow.pytorch
from darts.models import TFTModel

# 假设你已经训练好tft_model
tft_model = TFTModel(input_chunk_length=24, output_chunk_length=12)
tft_model.fit(train_series)

# 保存到本地或MLflow仓库
mlflow.pytorch.save_model(tft_model, "saved_tft_model")

加载模型

import mlflow.pytorch

loaded_tft = mlflow.pytorch.load_model("saved_tft_model")
# 正常使用预测功能
predictions = loaded_tft.predict(n=12, series=test_series)

方案2:使用Darts内置的MLflow集成

Darts的TFTModel自带log()方法,可以直接将模型日志到MLflow,无需手动指定flavor:

日志模型到MLflow

from darts.models import TFTModel
import mlflow

mlflow.start_run(run_name="tft_experiment")
tft_model = TFTModel(input_chunk_length=24, output_chunk_length=12)
tft_model.fit(train_series)
# 自动将模型、参数、指标等日志到当前MLflow run
tft_model.log()
mlflow.end_run()

从MLflow加载模型

from darts.models import TFTModel

# 通过run ID加载模型
loaded_tft = TFTModel.load_from_run(run_id="your_run_id_here")
predictions = loaded_tft.predict(n=12, series=test_series)

注意事项

  • 确保MLflow、Darts、PyTorch Lightning版本兼容,建议使用Darts官方推荐的依赖版本
  • 训练时如果使用了自定义回调或不可序列化的对象,需要提前处理(比如替换为可pickle的实现),避免保存失败

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 15:35:25