如何在MLflow中记录Darts的NBEATS模型并解决类型错误?
解决Darts NBEATS模型无法用mlflow.pytorch.log_model记录的问题
问题原因
Darts库的NBEATSModel是封装了PyTorch核心模型的高层类,并非原生的torch.nn.Module对象,而mlflow.pytorch.log_model仅接受原生PyTorch模块作为输入,因此直接传入Darts模型会触发TypeError。
解决方案
方案一:提取内部PyTorch核心模型记录
Darts的NBEATS模型内部通过model属性暴露了底层的PyTorch Module,可直接提取该对象用于MLflow记录,同时需要适配预测流程:
- 提取核心模型
# 假设nbeats_model是训练好的Darts NBEATSModel实例 pytorch_core_model = nbeats_model.model
- 记录模型并自定义预测逻辑
由于直接加载的PyTorch模型无法直接处理Darts的TimeSeries输入,需自定义预测函数适配格式:
import mlflow import torch import numpy as np from darts import TimeSeries def custom_predict(input_series, model): # 将TimeSeries转换为模型所需的张量 x_tensor = torch.tensor(input_series.values(copy=False).astype('float32')).unsqueeze(0) with torch.no_grad(): pred_tensor = model(x_tensor) # 将张量输出转回TimeSeries return TimeSeries.from_values(pred_tensor.squeeze(0).numpy()) # 生成输入示例用于模型签名 input_example = TimeSeries.from_values(np.array([[1.0], [2.0], [3.0]])) # 自动推断模型签名 signature = mlflow.models.infer_signature(input_example, custom_predict(input_example, pytorch_core_model)) # 记录模型 with mlflow.start_run(): mlflow.pytorch.log_model( pytorch_core_model, "nbeats_core_model", signature=signature, input_example=input_example, pip_requirements=["darts", "torch", "mlflow"] )
- 加载模型并预测
# 替换为你的run ID loaded_model = mlflow.pytorch.load_model("runs:/<your_run_id>/nbeats_core_model") test_series = TimeSeries.from_values(np.array([[4.0], [5.0], [6.0]])) prediction = custom_predict(test_series, loaded_model)
方案二:用MLflow PyFunc包装整个Darts模型
如果希望保留Darts模型的原生接口(如直接调用predict方法处理TimeSeries),可以用MLflow的PyFunc包装整个Darts模型:
- 定义PyFunc包装类
import mlflow.pyfunc from darts.models import NBEATSModel import pandas as pd class DartsNBEATSWrapper(mlflow.pyfunc.PythonModel): def load_context(self, context): # 加载保存的Darts模型 self.model = NBEATSModel.load(context.artifacts["saved_darts_model"]) def predict(self, context, model_input): # 将输入数据转换为Darts TimeSeries if isinstance(model_input, pd.DataFrame): input_series = TimeSeries.from_values(model_input.values) else: input_series = TimeSeries.from_values(model_input) # 执行预测(可根据需求调整预测步数n) return self.model.predict(n=len(input_series))
- 保存并记录模型
# 先将Darts模型保存到本地 nbeats_model.save("local_nbeats_model") # 记录PyFunc模型 with mlflow.start_run(): mlflow.pyfunc.log_model( "nbeats_pyfunc_model", python_model=DartsNBEATSWrapper(), artifacts={"saved_darts_model": "local_nbeats_model"}, pip_requirements=["darts", "mlflow", "torch", "pandas"] )
- 加载并预测
loaded_pyfunc_model = mlflow.pyfunc.load_model("runs:/<your_run_id>/nbeats_pyfunc_model") # 测试输入可以是DataFrame或numpy数组 test_input = pd.DataFrame([[7.0], [8.0], [9.0]]) prediction = loaded_pyfunc_model.predict(test_input)
内容的提问来源于stack exchange,提问作者Akshay Mitra
相关产品推荐
相关产品推荐

