MLflow无法记录Sklearn模型及下载模型运行报错问题求助
问题1:MLflow元数据记录失败警告
解决方案
- 统一MLflow版本:客户端与服务器版本不兼容是常见诱因,先检查当前版本:
升级至最新稳定版:mlflow --version
升级完成后重启MLflow服务再测试。pip install --upgrade mlflow - 验证Tracking URI配置:确认
sqlite:///mlflow.db路径正确,当前用户对该数据库文件拥有读写权限。可尝试删除旧的mlflow.db文件重新初始化,或改用绝对路径测试。 - 手动记录元数据/模型:若自动日志仍异常,改用手动记录兜底:
with mlflow.start_run(): # 记录指标 mlflow.log_metric("train_r2", baseline_train_r2) mlflow.log_metric("test_r2", baseline_test_r2) mlflow.log_metric("cv_r2", baseline_cv_r2) # 手动记录模型 mlflow.sklearn.log_model(lr_baseline, "baseline_model")
问题2:模型预测时输入特征不匹配
核心原因
训练时的X_train是经过独热编码后的数值特征,但预测时传入的df_submit_cleaned是原始分类特征,两者特征结构完全不一致。
解决方案
- 用Pipeline打包预处理+模型:这是最规范的解决方式,将特征工程与模型封装为一个整体,保存后预测时可直接传入原始数据:
后续加载模型预测时,直接传入原始数据即可,Pipeline会自动完成预处理步骤。from sklearn.pipeline import Pipeline from sklearn.compose import ColumnTransformer from sklearn.preprocessing import OneHotEncoder from sklearn.linear_model import LinearRegression # 假设分类特征列为cat_cols,数值特征列为num_cols preprocessor = ColumnTransformer( transformers=[ ('cat', OneHotEncoder(handle_unknown='ignore'), cat_cols), ('num', 'passthrough', num_cols) ]) # 构建Pipeline pipeline = Pipeline(steps=[ ('preprocessor', preprocessor), ('model', LinearRegression()) ]) # 训练Pipeline pipeline.fit(X_train, y_train) # 用MLflow记录整个Pipeline with mlflow.start_run(): mlflow.sklearn.log_model(pipeline, "pipeline_model") # 记录指标 mlflow.log_metric("train_r2", pipeline.score(X_train, y_train)) mlflow.log_metric("test_r2", pipeline.score(X_test, y_test)) - 添加模型签名(可选):记录模型时生成特征签名,提前检测输入不匹配问题:
from mlflow.models.signature import infer_signature signature = infer_signature(X_train, pipeline.predict(X_train)) mlflow.sklearn.log_model(pipeline, "pipeline_model", signature=signature) - 补救已保存的模型:若已单独保存模型,需同时保存预处理编码器,预测前先对输入数据做相同处理:
# 训练时保存编码器 import joblib joblib.dump(preprocessor, "preprocessor.joblib") # 预测时加载编码器与模型 preprocessor = joblib.load("preprocessor.joblib") model = mlflow.pyfunc.load_model(model_uri="./model") # 先预处理输入数据 X_submit_processed = preprocessor.transform(df_submit_cleaned) # 再执行预测 model.predict(X_submit_processed)
内容的提问来源于stack exchange,提问作者John Jam
相关产品推荐
相关产品推荐

