news 2026/9/20 15:15:34

Spark 中利用 MLflow 加载机器学习模型并对 pandas-on-Spark DataFrame 进行预测

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Spark 中利用 MLflow 加载机器学习模型并对 pandas-on-Spark DataFrame 进行预测
  • 大数据
  • 数据分析
  • 批处理
  • 流处理
  • 机器学习
  • 图计算

【免费下载链接】spark

Apache Spark - A unified analytics engine for large-scale data processing

项目地址:https://gitcode.com/gh_mirrors/sp/spark
点击查看免费下载

导读

本文围绕 Apache Spark 的 pandas-on-Spark(即pyspark.pandas)与 MLflow 的集成模块pyspark.pandas.mlflow展开,系统讲解如何把任意实现了 MLflowpyfuncflavor 的模型(scikit-learn、PyTorch、TensorFlow 等)加载为统一的预测器,并直接作用在分布式 pandas-on-Spark DataFrame 上完成大规模推理。读完本文,你将掌握load_modelPythonModelWrapper.predict的完整用法、返回类型推断机制、模型与 DataFrame 列合并的注意事项,以及该模块在源码层的实现原理,可直接在自己的 MLflow 模型 + Spark 分布式推理场景中落地。

本文对应仓库中的 API 参考文档为 python/docs/source/reference/pyspark.pandas/ml.rst,核心实现位于 python/pyspark/pandas/mlflow.py。

模块定位:让 MLflow 模型与 pandas-on-Spark 数据框无缝对接

pyspark.pandas.mlflow是 pandas-on-Spark 的机器学习工具模块,解决的核心问题是:训练阶段用 MLflow 统一管理的模型,如何在 Spark 分布式环境中直接对 pandas-on-Spark DataFrame 做批量预测

从 API 参考文档的定义看,该模块的前提条件非常明确:

任意 MLflow 模型,只要实现了 "pyfunc" flavor,即可用于 pandas-on-Spark DataFrame。绝大多数主流框架(scikit-learn、pytorch、tensorflow 等)都满足这一要求。

同时文档给出了一个关键约束:

使用本模块必须安装 MLflow 包(The MLflow package must be installed in order to use this module)。

也就是说,pyspark.pandas.mlflow本身是一个轻量适配层:它不依赖任何特定框架,而是统一依赖 MLflow 的pyfunc接口,把"任意框架训练出的模型"翻译为"可在 Spark 上分布式执行的 UDF"。这种设计让模型注册、加载、推理的链路与具体框架解耦,模型只要能在 MLflow 中保存和加载,就能接入 pandas-on-Spark。

模块对外只暴露两个符号(见 python/pyspark/pandas/mlflow.py 的__all__):

  • PythonModelWrapper:围绕 MLflow Python 对象模型的封装类,作为 pandas-on-Spark 上的预测器;
  • load_model:加载 MLflow 模型并返回PythonModelWrapper

核心 API:load_model 与 PythonModelWrapper.predict

load_model:加载模型入口

load_model的函数签名与参数语义如下:

def load_model( model_uri: str, predict_type: Union[str, type, Dtype] = "infer" ) -> PythonModelWrapper:
参数类型说明
model_uristr指向模型的 URI,支持 MLflow 支持的各种寻址方式(如runs:/<run_id>/modelmodels:/<name>/<version>、本地文件系统路径等),详见 MLflow 文档
predict_typePython 基础类型 / numpy 基础类型 / Spark 类型 /'infer'调用模型predict时期望的返回类型;指定'infer'时,包装器会尝试依据模型类型自动推断返回类型

返回值为PythonModelWrapper。该包装器遵循mlflow.pyfunc.PythonModel的接口约定。

PythonModelWrapper.predict:对两种数据框做预测

def predict(self, data: Union[DataFrame, pd.DataFrame]) -> Union[Series, pd.Series]:

predict会根据输入类型走两条完全不同的执行路径:

  1. 输入为 pandas DataFrame:直接调用底层pyfunc对象的predict,返回的是底层模型的原生输出(通常是 pandas Series 或 numpy 数组),适合小批量、单机场景;
  2. 输入为 pandas-on-Spark DataFrame:通过 Spark UDF 将模型映射到分布式数据上执行,返回 pandas-on-Spark Series,可在集群上对海量数据进行并行推理。

如果传入其他类型,则会抛出ValueError("unknown data type: ...")

源码视角:模型封装与分布式推理是如何实现的

PythonModelWrapper的实现体现了"惰性加载 + 类型推断 + UDF 包装"三个关键设计,全部代码见 python/pyspark/pandas/mlflow.py。

1. 惰性加载三个底层对象(lazy_property)

包装器内部通过lazy_property缓存三个对象,首次访问时才真正初始化:

  • _return_type:由predict_type提示转换成的 SparkDataType。逻辑上,当提示为"infer"或为空时,默认使用np.float64——对应默认的连续值预测场景。转换通过pyspark.pandas.typedef中的as_spark_type完成;
  • _model:调用mlflow.pyfunc.load_model(model_uri=...)得到底层模型对象;
  • _model_udf:调用mlflow.pyfunc.spark_udf(spark, model_uri=..., result_type=self._return_type)得到可直接挂载到 Spark 列上的 UDF,其中spark取自pyspark.pandas.utils.default_session()

源码中还留有明确的 TODO 注释:目前返回类型推断逻辑比较简单,仅覆盖"连续值预测"这一默认场景;后续可针对sklearn.Classifier(应返回整数或类别)以及 PyTorch / TensorFlow / Keras 模型依据输出类型做更智能的推断,但作者认为这部分更适合放在 MLflow 侧而非此处完成。这意味着当前版本对分类模型的predict_type需要用户显式指定,而不是完全依赖自动推断。

2. 分布式路径:struct + spark_udf + 内部结构替换

对 pandas-on-Spark DataFrame 的预测是模块的核心价值所在,其执行链路为:

s = struct(*data.columns) # 将整行特征打包成一个 struct 列 return_col = self._model_udf(s) # 对 struct 列应用 MLflow pyfunc UDF column_labels = [(col,) for col in data._internal.spark_frame.select(return_col).columns] internal = data._internal.copy( column_labels=column_labels, data_spark_columns=[return_col], data_fields=None ) return first_series(DataFrame(internal)) # 取回第一列作为 pandas-on-Spark Series

可以看到实现要点:

  • pyspark.sql.functions.struct把 DataFrame 的所有特征列打包为一个 struct 列,一次传给 UDF,避免逐列序列化;
  • spark_udf生成的 UDF 在 Spark 引擎内分布式执行,每个 executor 上的分区数据由 MLflow 加载的模型完成本地推理;
  • 预测结果作为一个新的 Spark 列(return_col),通过data._internal.copy替换列标签与数据列,构造出一个新的 pandas-on-Spark DataFrame,再用first_series取出唯一的预测列为 Series。

这种"列级 UDF + 内部元数据复制"的方式,让用户感受到的 API 与 pandas 原生体验一致,而底层计算已经被 Spark 调度到集群上。

3. 模块在 pandas-on-Spark 中的集成

该模块与 pandas-on-Spark 生态有两处明显的集成痕迹:

  • 在 python/pyspark/pandas/namespace.py 的依赖版本检测列表中,mlflowpysparkpandasnumpypyarrow等并列,说明它是该发行版预期可用的可选依赖之一;
  • 在 python/pyspark/pandas/usage_logging/init.py 中,mlflow模块与mlflow.PythonModelWrapper类被纳入使用日志统计范围(以try/except ImportError包裹,未安装 MLflow 时静默跳过)。

完整实战:scikit-learn 模型 + MLflow + pandas-on-Spark 推理

以下示例完整来自 load_model 的 docstring,是模块自带的 doctest 级可运行示例,可直接复制验证。

第一步:初始化 MLflow 环境

from mlflow.tracking import MlflowClient, set_tracking_uri import mlflow.sklearn from tempfile import mkdtemp d = mkdtemp("pandas_on_spark_mlflow") set_tracking_uri(f"sqlite:///{d}/mlflow.db") # 使用本地 sqlite 作为 tracking 后端 client = MlflowClient() exp_id = mlflow.create_experiment("my_experiment") exp = mlflow.set_experiment("my_experiment")

第二步:训练并记录 scikit-learn 线性回归模型

假设我们要学习的目标函数为y = log(2 + x)(以x1x2为特征):

import pandas as pd import numpy as np from sklearn.linear_model import LinearRegression train = pd.DataFrame({"x1": np.arange(8), "x2": np.arange(8)**2, "y": np.log(2 + np.arange(8))}) train_x = train[["x1", "x2"]] train_y = train[["y"]] with mlflow.start_run(): lr = LinearRegression() lr.fit(train_x, train_y) mlflow.sklearn.log_model(lr, "model")

第三步:加载模型并对 pandas-on-Spark DataFrame 预测

from pyspark.pandas.mlflow import load_model import pyspark.pandas as ps run_info = client.search_runs(exp_id)[-1].info model = load_model("runs:/{run_id}/model".format(run_id=run_info.run_id)) prediction_df = ps.DataFrame({"x1": [2.0], "x2": [4.0]}) prediction_df["prediction"] = model.predict(prediction_df) print(prediction_df)

输出(不同环境浮点精度可能有细微差异):

x1 x2 prediction 0 2.0 4.0 1.355551

第四步:对 pandas DataFrame 预测(单机路径)

同一个model对象也接受 pandas DataFrame,返回底层 pyfunc 的原生输出:

model.predict(prediction_df[["x1", "x2"]].to_pandas()) # array([[1.35555142]])

值得注意的是:同一个PythonModelWrapper同时支持 pandas 与 pandas-on-Spark 两种输入,这种双路径设计让用户在"本地调试用小数据(pandas)"与"集群推理用大数据(pandas-on-Spark)"之间切换时无需更换 API。

关键注意事项:预测结果列如何与原始 DataFrame 合并

官方文档在Notes一节明确指出了当前版本的一个重要限制:

目前,模型预测结果只能与现有 DataFrame 合并回去。其他列必须手动 join。

即:model.predict(features)返回的 Series 只能赋值回用于预测的那个特征子集 DataFrame,不能直接赋值给包含额外列的父 DataFrame。例如以下代码会失败(报错信息为数据框未对齐):

df = ps.DataFrame({"x1": [2.0], "x2": [3.0], "z": [-1]}) features = df[["x1", "x2"]] y = model.predict(features) features["y"] = y # 可用:预测列拼回 features 自身 df["y"] = y # 会失败:features 与 df 结构不对齐

官方给出的当前 workaround 是使用.merge(),以特征值为连接键把预测结果拼回原表:

features["y"] = y everything = df.merge(features, on=["x1", "x2"]) print(everything)

输出:

x1 x2 z y 0 2.0 3.0 -1 1.376932

在实际生产代码中,若x1x2并非唯一键,merge可能造成行数膨胀,建议结合实际键的唯一性选择连接列,或先对 DataFrame 进行行号/自增键标记再 merge。

predict_type 的选择策略

predict_typeload_model中唯一需要用户决策的参数,官方语义为:

  • 传 Python 基础类型、numpy 基础类型或 Spark 类型:包装器通过as_spark_type直接转换为 SparkDataType,作为 UDF 的result_type
  • 'infer'(默认值):包装器尝试自动推断,但当前实现只覆盖默认的连续值预测场景,即推断为np.float64

从源码可以得出的实践建议是:

  1. 回归模型 / 输出连续浮点的模型:直接使用默认predict_type="infer"即可;
  2. 分类模型(输出整数类别或概率):鉴于 mlflow.py 中的 TODO 注释表明分类推断尚未实现,建议显式传入对应的 Spark 类型(如IntegerTypeDoubleType)以确保spark_udfresult_type与模型输出一致,避免类型不匹配导致的运行期错误;
  3. 返回多列输出的模型:当前实现只取 UDF 结果的第一列(first_series),多输出场景需要自行评估是否符合预期。

适用场景与边界总结

适合的使用方式

  • 团队已用 MLflow 统一管理模型(tracking、registry),希望在同一套 Spark 作业里完成"读数据 → 分布式推理 → 结果入库";
  • 模型来自任何支持 pyfunc flavor 的框架,希望推理代码与框架无关;
  • 需要在小数据(pandas)上快速验证、再平滑迁移到大数据的 pandas-on-Spark 推理。

需要留意的边界

  • 必须先安装 MLflow(pip install mlflow),否则from pyspark.pandas.mlflow import load_model无法正常使用;
  • 预测列合并受限,跨列赋值需走.merge()
  • 返回类型推断以np.float64为默认兜底,分类/多输出场景应显式声明predict_type
  • 该模块面向批量离线/近线推理,非实时单条请求的低延迟服务场景。

进一步探索

  • 模块 API 参考文档:python/docs/source/reference/pyspark.pandas/ml.rst
  • 完整源码与内嵌 doctest 示例:python/pyspark/pandas/mlflow.py
  • 类型转换工具as_spark_type定义:python/pyspark/pandas/typedef/typehints.py
  • 依赖版本检测(含 mlflow):python/pyspark/pandas/namespace.py
  • 使用日志统计集成:python/pyspark/pandas/usage_logging/init.py
  • 大数据
  • 数据分析
  • 批处理
  • 流处理
  • 机器学习
  • 图计算

【免费下载链接】spark

Apache Spark - A unified analytics engine for large-scale data processing

项目地址:https://gitcode.com/gh_mirrors/sp/spark
点击查看免费下载

相关推荐

上一篇:REFramework项目中native布局重复问题的分析与解决
下一篇:5分钟上手:直播中实时显示键盘和手柄输入的免费神器

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 15:15:27

OfficeUtils 3.5.0 办公文档批量处理与PDF转换实战指南

简介&#xff1a;这是一款面向日常办公人群的绿色辅助工具箱&#xff0c;主要用于解决Office/WPS/PDF使用中的高频难题&#xff0c;如PDF转Word、PDF图片提取、Excel图片表格识别、多列组合排序、工作表合并以及从身份证号提取生日等。软件无需安装&#xff0c;解压即用&#x…

作者头像 李华
网站建设 2026/9/20 15:14:33

P104/P106无头显卡驱动魔改全攻略:从INF修改到代码43解决

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/20 15:11:35

文化遗产数字化保护的优势与实施路径:从三维扫描到数字资产管理

简介&#xff1a;《文化遗产数字化保护的优势与路径》是一份聚焦文化遗产数字化保护的系统论述文档&#xff0c;面向文博从业者、数字人文研究者及政策制定者&#xff0c;系统解答数字化手段如何提升保护的安全性与传播力&#xff0c;以及如何从制度、资金、平台、人才等层面落…

作者头像 李华
网站建设 2026/9/20 15:11:16

实时行情API选型:从延迟到数据质量,量化实盘避坑指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/20 15:08:43

VSCode远程开发配置指南:Codex AI助手在远程服务器上的部署与避坑

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华