Pyspark如何从dense vector概率列提取对应预测的最大概率到新列
PySpark 提取多分类预测结果对应概率的实现方法
这里提供两种可行方案,可根据你使用的PySpark版本选择:
方法1:自定义UDF实现(全版本兼容)
该方案不依赖高版本特性,所有PySpark版本均可使用:
- 先导入所需依赖
from pyspark.sql.functions import udf from pyspark.sql.types import DoubleType
- 定义UDF实现按索引取概率的逻辑
# 参数为概率向量prob、预测索引pred,返回对应位置的概率值 get_prediction_prob = udf(lambda prob, pred: float(prob[int(pred)]), DoubleType())
- 为DataFrame新增对应概率列
# 替换your_df为你实际的DataFrame变量名 result_df = your_df.withColumn("max_prob", get_prediction_prob("probability", "prediction"))
注:这里将prediction转int是因为模型输出的prediction默认为Double类型,而向量索引需要传入整数。
方法2:内置函数实现(PySpark 3.0+ 推荐)
PySpark 3.0开始提供内置vector_to_array函数,无需自定义UDF,执行效率更高:
- 导入所需依赖
from pyspark.ml.functions import vector_to_array from pyspark.sql.functions import col
- 直接转换取值
result_df = your_df.withColumn("prob_array", vector_to_array("probability")) \ .withColumn("max_prob", col("prob_array")[col("prediction").cast("int")]) \ .drop("prob_array") # 不需要中间数组列可直接删除
逻辑说明
RandomForestClassifier输出的probability是长度等于类别总数的DenseVector,每个位置的值对应该索引下类别的预测概率,prediction列的值本身就是概率最高的类别的索引,因此直接取该索引对应的值就是预测结果对应的最高概率。
内容的提问来源于stack exchange,提问作者Deb
相关产品推荐
相关产品推荐

