PySpark中如何从Row类型probability列提取第二个元素生成新列?
PySpark提取DenseVector列指定元素的正确方法
你原来用[:,1]的NumPy切片语法在PySpark里不适用,因为PySpark的Column对象不支持这种操作,得用Spark内置的列操作函数,下面是两种可行的方案:
方案一:使用
getItem()方法(0-based索引)
这个方法基于0开始计数,要提取第二个元素(对应索引1),代码如下:outp = outp.select("*", col('probability').getItem(1).alias('prob'))方案二:使用
element_at()函数(1-based索引,Spark 2.4+支持)
这个函数从1开始计数,提取第二个元素直接传2,代码更直观:from pyspark.sql.functions import element_at outp = outp.select("*", element_at(col('probability'), 2).alias('prob'))
两种方法都能正确提取DenseVector中的第二个概率分量,比如从DenseVector([0.99,0.01])中得到0.01。
内容的提问来源于stack exchange,提问作者user007
相关产品推荐
相关产品推荐

