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

Darts开发的NBEATS模型无法在MLflow中日志记录,报内存不足错误

问题解决方案

一、内存错误分析与解决

你遇到的RuntimeError: DefaultCPUAllocator: not enough memory错误,发生在训练结束后PyTorch Lightning尝试将GPU上的模型移回CPU的阶段。结合代码和日志,原因及解决方法如下:

1. 核心原因

  • 你的NBEATS模型参数规模达69.4M,训练后包含梯度数据的模型在移回CPU时,超出了当前可用的CPU内存容量。
  • 错误使用mlflow.keras.log_model记录PyTorch模型,提前占用额外内存,加剧了内存紧张。

2. 具体解决措施

  • 禁用训练后自动移回CPU:在pl_trainer_kwargs中添加参数,阻止模型自动从GPU迁移到CPU:
    pl_trainer_kwargs={
        "accelerator": "gpu",
        "devices": [0],
        "enable_model_summary": False
    }
    
  • 手动清理内存:训练完成后释放GPU和CPU内存:
    model.fit(darts_y_train, past_covariates=darts_x_train, epochs=1)
    # 清理GPU缓存与冗余内存
    torch.cuda.empty_cache()
    import gc
    gc.collect()
    
  • 缩小模型规模:适当减小input_chunk_length,或调整NBEATS的stacks参数降低整体内存占用:
    nbeats_model = NBEATSModel(
        input_chunk_length=336,  # 从672减半
        output_chunk_length=output_chunk_length,
        n_epochs=1,
        random_state=42,
        stacks=10,  # 默认30,减少后降低参数量
        pl_trainer_kwargs={...}
    )
    

二、正确将NBEATS模型记录到MLflow

代码中错误使用mlflow.keras.log_model记录PyTorch架构的NBEATS模型,导致兼容性问题。正确做法是使用mlflow.pytorch.log_model,且必须在模型训练完成后执行日志记录:

修改后的mlflow_run函数:

def mlflow_run(run_name="nbeats_model_run"):
    with mlflow.start_run(run_name=run_name, nested=True) as run:
        model = baseline_model()
        # 先完成模型训练
        model.fit(darts_y_train, past_covariates=darts_x_train, epochs=1)
        # 训练后清理内存
        torch.cuda.empty_cache()
        # 使用PyTorch专属方法记录模型
        mlflow.pytorch.log_model(model=model, artifact_path="model")
        # 执行预测
        model.predict(n=1344, series=test_series, past_covariates=df_x)
        run_id = run.info.run_uuid
        exp_id = run.info.experiment_id

    return exp_id, run_id

三、NBEATS+MLflow的参考案例

Darts社区已有成功实践:

  • Darts官方文档的实验跟踪章节包含MLflow集成的通用示例,可直接适配NBEATS模型;
  • Darts的GitHub讨论区中,有用户分享过完整的NBEATS+MLflow训练日志流程,核心就是用mlflow.pytorch.log_model完成模型记录。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 15:50:39