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

已获取MLflow的run_id,如何正确加载对应的PyTorch模型

解决MLflow加载PyTorch模型的路径问题

核心问题原因

你调用mlflow.pytorch.load_model(f"runs:/{best_run_id}")失败,是因为MLflow的run URI需要指定模型在run artifacts中的相对路径。一个MLflow run可以保存多个文件(比如模型、数据、日志),你必须明确告诉MLflow要加载哪个文件夹下的模型文件——这个路径就是你调用mlflow.pytorch.log_model()时传入的第二个参数(比如mlflow.pytorch.log_model(model, "my_model")里的"my_model")。

解决方法

方法1:直接拼接已知的模型保存路径

如果你记得保存模型时指定的路径(比如当时用的是mlflow.pytorch.log_model(model, "model")),直接把路径拼到run_id后面即可:

best_run_id = get_best_run_id("<my_experiment_name>")
# 假设保存模型时的路径是"model"
model = mlflow.pytorch.load_model(f"runs:/{best_run_id}/model")

方法2:通过代码自动获取模型的artifact路径

如果不确定路径,可以通过mlflow.get_run()获取run的artifact信息,自动找到模型路径:

import mlflow

def get_best_run_id(experiment_name, metric_to_sort_by='f1_score'):
    # 简化实验ID获取逻辑
    experiment = mlflow.get_experiment_by_name(experiment_name)
    if not experiment:
        raise ValueError(f"找不到名为{experiment_name}的实验")
    runs_df = mlflow.search_runs(experiment.experiment_id)
    # 过滤已完成的run,保留参数和指标列
    runs_df_filtered = runs_df[
        (runs_df['status'] == 'FINISHED')
    ].filter(regex='^params\.|^metrics\.').dropna()
    # 按指定指标降序排序
    runs_df_sorted = runs_df_filtered.sort_values(f"metrics.{metric_to_sort_by}", ascending=False)
    if runs_df_sorted.empty:
        raise ValueError("没有符合条件的已完成run")
    # 关联原run数据获取run_id
    best_run_id = runs_df.loc[runs_df_sorted.index[0], 'run_id']
    return best_run_id

best_run_id = get_best_run_id("<my_experiment_name>")
# 获取run的artifact信息
run = mlflow.get_run(best_run_id)
# 查找模型类型的artifact(PyTorch模型会有MLmodel文件)
model_artifact_path = None
for artifact in run.data.artifacts:
    # 检查artifact是否包含MLmodel文件(MLflow模型的标识)
    if artifact.path.endswith("MLmodel"):
        # 取MLmodel文件所在的文件夹路径
        model_artifact_path = "/".join(artifact.path.split("/")[:-1])
        break

if model_artifact_path:
    model = mlflow.pytorch.load_model(f"runs:/{best_run_id}/{model_artifact_path}")
else:
    raise ValueError(f"run {best_run_id}中未找到PyTorch模型")

方法3:直接从Model Registry加载(更推荐)

既然你提到了MLflow Model Registry,如果你已经把模型注册到了Registry,完全可以不用通过run_id加载,直接用模型名称和版本/阶段加载:

# 加载指定版本的模型
model = mlflow.pytorch.load_model("models:/<你的模型名称>/1")
# 或者加载指定阶段的模型(比如Staging/Production)
model = mlflow.pytorch.load_model("models:/<你的模型名称>/Production")

额外提示

  • 保存模型时,务必记录好log_model()的路径参数,这是加载模型的关键。
  • 可以通过MLflow UI查看对应run的Artifacts标签页,直观看到模型的保存路径,直接复制使用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 17:05:18