Spark中不使用collect_list/collect_set实现分组拼接字符串
解决Spark DataFrame分组拼接字符串(无需collect_list/collect_set)
嘿,刚好碰到过类似需求,给你两种靠谱的解决办法,都不用你说的那两个函数:
方法1:用Spark内置的string_agg函数(推荐,Spark 2.4+适用)
Spark 2.4及以上版本提供了string_agg这个专门的字符串聚合函数,直接就能实现分组拼接,完全不需要借助collect_list或者collect_set。
代码示例
首先先还原你的DataFrame:
from pyspark.sql import SparkSession from pyspark.sql.functions import string_agg spark = SparkSession.builder.appName("StringConcatDemo").getOrCreate() # 你的原始数据 data = [ ("9d71445e-ee5d-4d37-bfb7-02f6e6eacd9d", "Friday 0 0.9604490986400536"), ("9d71445e-ee5d-4d37-bfb7-02f6e6eacd9d", "Friday 1 0.8109076852795446"), ("9d71445e-ee5d-4d37-bfb7-02f6e6eacd9d", "Friday 2 0.7282039568471731"), ("9d71445e-ee5d-4d37-bfb7-02f6e6eacd9d", "Friday 3 0.5335418350493728") ] df = spark.createDataFrame(data, ["MeteVarID", "Conc"])
然后执行分组拼接:
# 按MeteVarID分组,拼接Conc列,用逗号加空格分隔 result_df = df.groupBy("MeteVarID").agg( string_agg("Conc", ", ").alias("Conc_combined") ) # 查看结果 result_df.show(truncate=False)
输出结果
+------------------------------------+----------------------------------------------------------------------------------------------------+ |MeteVarID |Conc_combined | +------------------------------------+----------------------------------------------------------------------------------------------------+ |9d71445e-ee5d-4d37-bfb7-02f6e6eacd9d|Friday 0 0.9604490986400536, Friday 1 0.8109076852795446, Friday 2 0.7282039568471731, Friday 3 0.5335418350493728| +------------------------------------+----------------------------------------------------------------------------------------------------+
这个函数的好处就是简洁高效,是Spark官方为字符串聚合场景提供的原生解决方案。
方法2:自定义UDAF(适用于Spark 2.4以下版本)
如果你的Spark版本低于2.4,没有string_agg,那可以自己写一个用户自定义聚合函数(UDAF)来实现拼接逻辑,同样不需要用到collect_list或collect_set。
代码示例
from pyspark.sql.types import StringType, StructType, StructField from pyspark.sql.udf import UserDefinedAggregateFunction class StringConcatUDAF(UserDefinedAggregateFunction): # 定义输入数据的类型(这里是单个字符串列) def inputSchema(self): return StructType([StructField("value", StringType())]) # 定义缓冲区的类型(用来存储中间拼接的结果) def bufferSchema(self): return StructType([StructField("concatenated_str", StringType())]) # 定义输出结果的类型 def dataType(self): return StringType() # 标记函数是否是确定性的(相同输入总是返回相同输出) def deterministic(self): return True # 初始化缓冲区:刚开始是空字符串 def initialize(self, buffer): buffer[0] = "" # 每处理一条数据时更新缓冲区:把当前字符串拼接到已有结果后面 def update(self, buffer, input): if buffer[0] == "": buffer[0] = input[0] else: buffer[0] = f"{buffer[0]}, {input[0]}" # 合并多个分区的缓冲区结果 def merge(self, buffer1, buffer2): if buffer1[0] == "": buffer1[0] = buffer2[0] elif buffer2[0] != "": buffer1[0] = f"{buffer1[0]}, {buffer2[0]}" # 生成最终的输出结果 def evaluate(self, buffer): return buffer[0] # 注册这个自定义UDAF concat_udaf = StringConcatUDAF() # 使用UDAF进行分组拼接 result_df = df.groupBy("MeteVarID").agg( concat_udaf("Conc").alias("Conc_combined") ) result_df.show(truncate=False)
这个UDAF的逻辑很直白:初始化一个空字符串,每收到一条数据就把它拼接到缓冲区里,最后把缓冲区的结果作为分组后的拼接值返回。
内容的提问来源于stack exchange,提问作者Long Time no see
相关产品推荐
相关产品推荐

