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

如何在PySpark中用Scikit-learn训练的XGBoost模型推理?非UDF方案有哪些?

在PySpark中使用Scikit-learn训练的XGBoost模型推理的方案

方案1:转换为PMML格式推理

PMML是跨平台的模型交换格式,支持多数机器学习模型(包括XGBoost)。你可以先将Sklearn训练的XGBoost模型导出为PMML,再用PySpark加载PMML进行批量推理。

步骤:

  • 安装依赖:sklearn2pmml 和 pyspark-pmml
  • 导出模型为PMML:
from sklearn2pmml import PMMLPipeline, sklearn2pmml
import xgboost as xgb
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split

# 示例训练模型
X, y = load_iris(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)
xgb_model = xgb.XGBClassifier()
xgb_model.fit(X_train, y_train)

# 用PMMLPipeline包装模型并导出
pipeline = PMMLPipeline([("classifier", xgb_model)])
sklearn2pmml(pipeline, "xgb_iris_model.pmml")
  • PySpark中加载PMML推理:
from pyspark.sql import SparkSession
from pyspark_pmml import PMMLModel

spark = SparkSession.builder.appName("XGBoostPMMLInference").getOrCreate()

# 加载数据集(示例用iris数据)
df = spark.createDataFrame([tuple(row) for row in X_test.tolist()], ["sepal_length", "sepal_width", "petal_length", "petal_width"])

# 加载PMML模型
pmml_model = PMMLModel.fromFile("xgb_iris_model.pmml")
result_df = pmml_model.transform(df)

result_df.show()

注意事项:

  • 确保sklearn2pmml版本与你的Sklearn、XGBoost版本兼容,避免导出失败
  • 复杂的预处理逻辑需要放到PMMLPipeline中,否则推理时会缺失步骤

方案2:转换为ONNX格式推理

ONNX是跨框架的模型标准,支持模型在不同平台间迁移。可以将Sklearn训练的XGBoost模型转换为ONNX,再用PySpark结合ONNX Runtime进行分布式推理。

步骤:

  • 安装依赖:onnx, skl2onnx, onnxruntime, pyspark
  • 转换模型为ONNX:
from skl2onnx import convert_sklearn
from skl2onnx.common.data_types import FloatTensorType
import onnx

# 定义输入张量类型(根据特征数量调整)
initial_type = [("float_input", FloatTensorType([None, 4]))]
onnx_model = convert_sklearn(xgb_model, initial_types=initial_type)
onnx.save(onnx_model, "xgb_iris_model.onnx")
  • PySpark中用ONNX Runtime推理(推荐用Pandas UDF,比普通UDF高效):
import onnxruntime as rt
import pandas as pd
from pyspark.sql.functions import pandas_udf, col

# 加载ONNX模型到内存(每个Executor会加载一次)
sess = rt.InferenceSession("xgb_iris_model.onnx")
input_name = sess.get_inputs()[0].name
label_name = sess.get_outputs()[0].name

# 定义Pandas UDF
@pandas_udf("int")
def predict_udf(df: pd.DataFrame) -> pd.Series:
    inputs = df.values.astype("float32")
    pred = sess.run([label_name], {input_name: inputs})[0]
    return pd.Series(pred)

# 应用UDF到数据集
result_df = df.withColumn("prediction", predict_udf(col("sepal_length"), col("sepal_width"), col("petal_length"), col("petal_width")))
result_df.show()

注意事项:

  • 转换ONNX时要确保输入特征的数量、类型与训练时一致
  • 如果模型有预处理步骤(如标准化),需要在PySpark中先完成预处理,或者将预处理逻辑也转换为ONNX算子

方案3:使用Spark MLlib的XGBoost接口加载兼容模型

如果你训练的XGBoost模型是用xgboost.sklearn接口训练的,可以尝试将模型保存为XGBoost原生格式,再用Spark MLlib的XGBoost加载器加载,直接进行分布式推理。

步骤:

  • 保存Sklearn训练的XGBoost模型为原生格式:
xgb_model.save_model("xgb_iris_model.model")
  • PySpark中用MLlib加载模型并推理:
from pyspark.ml.feature import VectorAssembler
from pyspark.ml.classification import XGBoostClassifier

# 将特征列合并为向量列(Spark MLlib要求输入为Vector类型)
assembler = VectorAssembler(inputCols=["sepal_length", "sepal_width", "petal_length", "petal_width"], outputCol="features")
df_vector = assembler.transform(df)

# 加载XGBoost模型
spark_xgb_model = XGBoostClassifier.load("xgb_iris_model.model")
result_df = spark_xgb_model.transform(df_vector)

result_df.select("sepal_length", "sepal_width", "petal_length", "petal_width", "prediction").show()

注意事项:

  • 此方法要求Spark MLlib的XGBoost版本与训练时的XGBoost版本兼容,版本差异可能导致加载失败
  • 仅支持XGBoost原生模型格式,且模型的参数设置需要符合Spark MLlib的要求

方案4:Pandas UDF(Vectorized UDF)替代普通UDF

如果必须用UDF,Pandas UDF比普通Python UDF效率高很多——它批量处理数据,减少了Python与JVM之间的序列化开销。

示例代码:

import pandas as pd
from pyspark.sql.functions import pandas_udf
import xgboost as xgb

# 加载训练好的Sklearn XGBoost模型
loaded_model = xgb.XGBClassifier()
loaded_model.load_model("xgb_iris_model.model")

# 定义Pandas UDF
@pandas_udf("int")
def predict_batch(df: pd.DataFrame) -> pd.Series:
    # df是包含所有特征列的DataFrame
    predictions = loaded_model.predict(df)
    return pd.Series(predictions)

# 应用UDF
result_df = df.withColumn("prediction", predict_batch(*df.columns))
result_df.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 21:13:26