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

如何在PySpark DataFrame上应用需分类列的LGBM模型做预测

解决PySpark DataFrame上LGBM模型预测的问题

针对你的场景,核心问题是PySpark DataFrame的字符串分类特征与LGBM模型输入的类型对齐,以及如何在分布式环境中调用本地训练的LGBM模型。以下是三种可行的解决方案,按需选择:


方案一:小数据量场景 - 转Pandas DataFrame预测

如果数据量不大(百万条以内),直接转Pandas是最简单的方式:

  1. 将PySpark DataFrame转为Pandas

    df_pd = spark_df.toPandas()
    
  2. 对齐特征类型
    根据模型训练时的特征类型处理:

    • 若训练时RESIDENCE_GROUP和SR_BRANDMODEL是数值类型(int),把PySpark中对应的字符串列转成int:
      df_pd['RESIDENCE_GROUP'] = df_pd['RESIDENCE_GROUP'].astype(int)
      df_pd['SR_BRANDMODEL'] = df_pd['SR_BRANDMODEL'].astype(int)
      
    • DISTRITO保持字符串类型即可。
  3. 确保特征顺序一致
    必须和模型训练时的特征顺序完全匹配,可通过model.feature_name_获取训练时的特征列表:

    import pickle
    # 加载pickle模型
    with open('lgbm_model.pkl', 'rb') as f:
        model = pickle.load(f)
    
    # 按训练时的特征顺序提取数据
    X = df_pd[model.feature_name_]
    
  4. 预测并转回PySpark

    df_pd['prediction'] = model.predict(X)
    result_spark = spark.createDataFrame(df_pd)
    

方案二:中大数据量场景 - PySpark UDF分布式预测

如果数据量较大但未到超大规模,用PySpark UDF结合广播变量实现分布式预测:

  1. 广播模型到所有Executor
    避免每个Task重复加载模型,提升性能:

    import pickle
    from pyspark.sql import functions as F
    from pyspark.sql.types import FloatType
    
    # 加载模型
    with open('lgbm_model.pkl', 'rb') as f:
        model = pickle.load(f)
    
    # 广播模型(每个Executor仅加载一次)
    broadcast_model = spark.sparkContext.broadcast(model)
    
  2. 定义特征顺序
    必须和训练时的特征顺序一致:

    # 替换为你的15个特征的实际顺序,需与model.feature_name_一致
    feature_order = ['feat1', 'feat2', ..., 'RESIDENCE_GROUP', 'DISTRITO', 'SR_BRANDMODEL']
    
  3. 定义预测UDF
    在UDF中处理每行特征并调用模型:

    @F.udf(returnType=FloatType())
    def predict_lgbm(*features):
        import pandas as pd
        # 构造单行数据,保持特征顺序
        df_row = pd.DataFrame([list(features)], columns=feature_order)
        # 对齐特征类型(根据训练时的类型调整)
        df_row['RESIDENCE_GROUP'] = df_row['RESIDENCE_GROUP'].astype(int)
        df_row['SR_BRANDMODEL'] = df_row['SR_BRANDMODEL'].astype(int)
        # 执行预测
        pred = broadcast_model.value.predict(df_row)[0]
        return float(pred)
    
  4. 调用UDF生成预测结果

    result_df = spark_df.withColumn('prediction', predict_lgbm(*feature_order))
    

方案三:超大数据量场景 - 转ONNX格式预测

对于超大规模数据,将LGBM模型转成ONNX格式,利用ONNX Runtime的高性能分布式预测:

第一步:将LGBM模型转为ONNX格式

需要安装onnxmltools和skl2onnx,注意版本兼容(建议与你的LGBM版本匹配):

import onnxmltools
from onnxmltools.convert.common.data_types import FloatTensorType, StringTensorType

# 定义输入类型:根据你的15个特征类型定义,假设前12个是数值型,后3个是分类字符串
input_types = []
# 数值特征
for feat_name in model.feature_name_[:12]:
    input_types.append((feat_name, FloatTensorType([None, 1])))
# 分类特征
input_types.append(('RESIDENCE_GROUP', StringTensorType([None, 1])))
input_types.append(('DISTRITO', StringTensorType([None, 1])))
input_types.append(('SR_BRANDMODEL', StringTensorType([None, 1])))

# 转换模型
onnx_model = onnxmltools.convert_lightgbm(model, initial_types=input_types, target_opset=12)

# 保存ONNX模型
with open('lgbm_model.onnx', 'wb') as f:
    f.write(onnx_model.SerializeToString())

第二步:PySpark中用ONNX Runtime预测

import onnxruntime as rt
from pyspark.sql import functions as F
from pyspark.sql.types import FloatType

# 广播ONNX模型路径
broadcast_model_path = spark.sparkContext.broadcast('lgbm_model.onnx')

# 定义预测UDF
@F.udf(returnType=FloatType())
def predict_onnx(*features):
    # 初始化ONNX会话(每个Executor第一次调用时创建)
    sess = rt.InferenceSession(broadcast_model_path.value)
    # 构造输入字典
    input_names = [inp.name for inp in sess.get_inputs()]
    input_dict = {}
    for name, val in zip(input_names, features):
        if name in ['RESIDENCE_GROUP', 'DISTRITO', 'SR_BRANDMODEL']:
            # 分类特征转字符串数组
            input_dict[name] = [[val]]
        else:
            # 数值特征转float数组
            input_dict[name] = [[float(val)]]
    # 执行预测
    pred = sess.run(None, input_dict)[0][0]
    return float(pred)

# 生成结果
result_df = spark_df.withColumn('prediction', predict_onnx(*model.feature_name_))

关键注意事项

  1. 特征顺序绝对不能错:LGBM模型对输入特征的顺序敏感,必须和训练时完全一致,可通过model.feature_name_获取训练时的特征列表。
  2. 分类特征类型必须对齐:如果训练时RESIDENCE_GROUP和SR_BRANDMODEL是数值类型,必须将PySpark中的字符串列转为对应数值类型;如果训练时是字符串类型,则保持字符串即可。
  3. 性能权衡:小数据用方案一,中数据用方案二,超大数据用方案三。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 08:28:10