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

如何在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记录,同时需要适配预测流程:

  1. 提取核心模型
# 假设nbeats_model是训练好的Darts NBEATSModel实例
pytorch_core_model = nbeats_model.model
  1. 记录模型并自定义预测逻辑
    由于直接加载的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"]
    )
  1. 加载模型并预测
# 替换为你的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模型:

  1. 定义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))
  1. 保存并记录模型
# 先将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"]
    )
  1. 加载并预测
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 20:55:34