如何在MLFlow模型中保存自定义额外元数据并实现便捷访问
MLFlow 存储模型自定义元数据解决方案
需求背景
使用MLFlow保存模型时,因存在自定义预测流水线拉取数据的场景,需要向模型中存入额外元数据。
前期尝试的问题
- 使用自定义signature类存储元数据:可正常写入MLModel的YAML文件,但从MLFlow registry加载模型时,signature访问效率极低,不符合使用要求。
对应代码:mlflow.sklearn.log_model(model, "model", signature = signature) - 直接向log_model传入额外字典:元数据会被写入conda.yaml文件中,不符合配置文件分工要求。
对应代码:mlflow.sklearn.log_model(model, "model", {"metadata1":"value1", "metadata2":"value2"})
最终方案
无需自定义MLFlow flavour或者自定义Model继承类,MLFlow原生log_model方法提供了metadata参数专门用于存储自定义元数据,元数据会自动写入MLModel配置文件,加载时可直接访问。
注意:该特性在MLFlow 1.27.0及以上版本支持,版本较低的环境请先升级MLFlow
修改后的存储代码
import mlflow import numpy as np import pandas as pd from sklearn.linear_model import LogisticRegression metadata_dic = {"metadata1": "value1", "metadata2": "value2"} X = np.array([[-2, -1, 0, 1, 2, 1],[-2, -1, 0, 1, 2, 1]]).T y = np.array([0, 0, 1, 1, 1, 0]) X = pd.DataFrame(X, columns=["X1", "X2"]) y = pd.DataFrame(y, columns=["y"]) model = LogisticRegression() model.fit(X, y) # 传入metadata参数存储自定义元数据 mlflow.sklearn.log_model( model, "model", metadata=metadata_dic )
加载模型读取元数据代码
import mlflow.pyfunc # 替换为你的模型路径,支持本地路径、runs路径、registry路径 loaded_model = mlflow.pyfunc.load_model("runs:/<运行ID>/model") # 读取自定义元数据 print(loaded_model.metadata.metadata["metadata1"]) # 输出:value1
内容的提问来源于stack exchange,提问作者Angelo
相关产品推荐
相关产品推荐

