基于PySpark DataFrame用XGBooster预测报错的解决方法
问题描述
我有一个包含通用列、特征列和标签列的预测输入数据集,以及一个xgb.Booster类型的XGBoost模型,尝试使用以下代码基于该模型执行预测:
from pyspark.sql.functions import col, DataFrame, udf, array from pyspark.sql.types import FloatType feature_cols = ["feature1", "feature2", "feature3"] label_col = "label" all_cols = ["id1", "id2"] + feature_cols + label_col pred_col_name = "pred" prediction_path = "/input/path" prediction_input = spark.read.parquet(prediction_path).select(all_cols) for column in feature_cols: prediction_input = prediction_input.withColumn(column, col(column).cast("float")) def predict_udf(*features): dmatrix = xgb.DMatrix(list(features)) return float(model.predict(dmatrix)[0]) predict_udf_spark = udf(predict_udf, FloatType()) mask_labeled = prediction_input.filter(col(label_col) != 0) if mask_labeled.count() > 0: prediction_labeled = mask_labeled.withColumn(pred_col_name, predict_udf_spark(*[mask_labeled[col] for col in feature_cols])) prediction_labeled.select(*output_cols).repartition(1000).write.csv(output_path, header=False, mode="append")
报错信息
运行代码时触发以下错误:
jc = sc._jvm.functions.array(_to_seq(sc, cols, _to_java_column)) AttributeError: 'NoneType' object has no attribute '_jvm'
解决方案
核心原因
错误源于UDF调用时的列传递方式错误:mask_labeled[col]的引用方式导致Spark在处理列时出现SparkContext(sc)为None的异常,本质是列对象的传递逻辑不符合Spark UDF的要求。
修复步骤
修正UDF的列传递逻辑
替换错误的列引用方式,直接使用col()函数或者列名字符串传递特征列:# 方式1:使用col()函数包装列名 prediction_labeled = mask_labeled.withColumn(pred_col_name, predict_udf_spark(*[col(c) for c in feature_cols])) # 方式2:直接传递列名字符串(Spark UDF支持直接解析) prediction_labeled = mask_labeled.withColumn(pred_col_name, predict_udf_spark(*feature_cols))额外问题修复与优化
- 修正
all_cols的定义:label_col是字符串,直接与列表相加会报错,改为:all_cols = ["id1", "id2"] + feature_cols + [label_col] - 优化UDF性能:避免每次调用UDF都创建
xgb.DMatrix,可以先将特征合并为数组列再传入UDF:from pyspark.sql.functions import array # 合并特征列为数组 prediction_input = prediction_input.withColumn("features_array", array(*feature_cols)) # 修改UDF接收数组参数 def predict_udf(features): dmatrix = xgb.DMatrix([features]) return float(model.predict(dmatrix)[0]) predict_udf_spark = udf(predict_udf, FloatType()) # 调用UDF时传入数组列 prediction_labeled = mask_labeled.withColumn(pred_col_name, predict_udf_spark(col("features_array"))) - 确保
model变量在UDF的作用域内,且在分布式环境中支持序列化(可使用cloudpickle等工具序列化模型,确保能传递到各个Executor节点)
- 修正
内容的提问来源于stack exchange,提问作者Utkarsh Roy
相关产品推荐
相关产品推荐

