PySpark中使用collect_set后如何去除双重括号?
解决PySpark分组后嵌套列表扁平化问题
问题场景
现有如下PySpark DataFrame,其中codes字段是字符串格式的列表:
DF = [('1', '[132]'), ('1', '[184, 88]'), ('2', '[55]'), ('2', '[123,33]'),] DF = spark.sparkContext.parallelize(DF).toDF(['id', 'codes'])
执行分组聚合后:
DF.groupBy("id").agg(F.collect_set("codes").alias("codes_concat")).show(4)
得到嵌套列表格式的结果:
+---+------------------+ | id| codes_concat| +---+------------------+ | 1|[[184, 88], [132]]| | 2| [[123,33], [55]]| +---+------------------+
需要将codes_concat转换为单层扁平化列表,目标输出:
+---+------------------+ | id| codes_concat| +---+------------------+ | 1| [184, 88, 132] | | 2| [123,33, 55] | +---+------------------+
解决方案
方法一:先解析字符串为数组,再聚合扁平化
先把字符串格式的codes转换为PySpark数组类型,再分组聚合后直接扁平化,这种方式更高效:
import pyspark.sql.functions as F # 1. 将字符串列表解析为整数数组 df_parsed = DF.withColumn( "codes_array", # 去掉首尾方括号,按逗号(含可选空格)分割,转整数数组 F.split(F.regexp_replace("codes", r"^\[|\]$", ""), ",\s*").cast("array<int>") ) # 2. 分组聚合后扁平化数组,如需去重则添加array_distinct df_result = df_parsed.groupBy("id").agg( F.flatten(F.collect_set("codes_array")).alias("codes_concat") # 若要确保最终列表元素唯一,使用: # F.array_distinct(F.flatten(F.collect_set("codes_array"))).alias("codes_concat") ) df_result.show(truncate=False)
方法二:对已分组的嵌套列表做后处理
如果已经得到了分组后的嵌套列表结果,可以直接对codes_concat字段做转换:
import pyspark.sql.functions as F # 先执行原分组操作 df_grouped = DF.groupBy("id").agg(F.collect_set("codes").alias("codes_concat")) # 遍历嵌套的字符串元素,逐个解析为数组后扁平化 df_result = df_grouped.withColumn( "codes_concat", F.flatten( F.transform( "codes_concat", lambda s: F.split(F.regexp_replace(s, r"^\[|\]$", ""), ",\s*").cast("array<int>") ) ) ) df_result.show(truncate=False)
内容的提问来源于stack exchange,提问作者user19495470
相关产品推荐
相关产品推荐

