如何将PySpark DataFrame指定列的所有行提取为array类型容器
PySpark提取指定列全量数据为数组的实现方案
你之前使用的my_array = df.select(df['my_col'])返回的是仅包含my_col列的PySpark DataFrame对象,并非本地数组结构,所以不符合需求,可通过以下方案实现需求:
方案1:直接聚合后拉取(适用于可直接序列化的UDT类型、小数据集场景)
调用PySpark内置的collect_list函数,先将全量行的指定列聚合为单行列的数组结构,再拉取到本地即可:
from pyspark.sql import functions as F # 聚合指定列到单行列的数组 agg_df = df.agg(F.collect_list("my_col").alias("col_arr")) # 拉取结果到本地,得到array类型容器 my_array = agg_df.first()["col_arr"]
方案2:UDT转基础类型后聚合(适用于序列化存在兼容性问题的UDT场景)
如果直接聚合UDT类型列后拉取得到的结果无法正常解析,可以先通过UDF将UDT转换为Python原生支持的基础类型(如列表、数值等),再执行聚合操作。
以常见的VectorUDT类型为例,代码示例:
from pyspark.sql import functions as F from pyspark.sql.types import ArrayType, DoubleType from pyspark.ml.linalg import Vector # 定义UDF将VectorUDT转为列表 vec_to_list_udf = F.udf(lambda vec: vec.toArray().tolist(), ArrayType(DoubleType())) # 新增转换后的基础类型列 df = df.withColumn("my_col_base", vec_to_list_udf("my_col")) # 聚合后拉取得到数组 my_array = df.agg(F.collect_list("my_col_base")).first()[0]
注意事项
- 上述方案都会将全量列数据拉取到Driver节点内存中,仅适用于小数据集场景。如果数据量超过Driver内存上限,建议直接在Spark分布式侧完成后续计算,无需拉取到本地转成数组
collect_list默认会跳过列中的null值,如果需要保留null值,可以搭配条件判断函数处理:F.collect_list(F.when(F.col("my_col").isNotNull(), F.col("my_col")).otherwise(F.lit(None)))
内容的提问来源于stack exchange,提问作者Morello
相关产品推荐
相关产品推荐

