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

Pyspark如何从dense vector概率列提取对应预测的最大概率到新列

PySpark 提取多分类预测结果对应概率的实现方法

这里提供两种可行方案,可根据你使用的PySpark版本选择:


方法1:自定义UDF实现(全版本兼容)

该方案不依赖高版本特性,所有PySpark版本均可使用:

  1. 先导入所需依赖
from pyspark.sql.functions import udf
from pyspark.sql.types import DoubleType
  1. 定义UDF实现按索引取概率的逻辑
# 参数为概率向量prob、预测索引pred,返回对应位置的概率值
get_prediction_prob = udf(lambda prob, pred: float(prob[int(pred)]), DoubleType())
  1. 为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,执行效率更高:

  1. 导入所需依赖
from pyspark.ml.functions import vector_to_array
from pyspark.sql.functions import col
  1. 直接转换取值
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 01:45:03