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

