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

