如何用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
相关产品推荐
相关产品推荐

