Spark DataFrame分组统计问题:按key1分组统计key2数量并新增列
解决Spark DataFrame按key1分组统计key2计数并添加到原数据的问题
你的需求是给原DataFrame新增一列,记录每条数据对应的key1分组下key2的出现次数,原代码的问题在于连续使用groupBy的方式不对,而且无法将统计结果关联回原数据。下面提供两种简洁的实现方法:
方法一:使用窗口函数(推荐)
窗口函数可以直接在原DataFrame上计算分组统计值,不需要额外的join操作,代码更简洁高效:
import org.apache.spark.sql.expressions.Window import org.apache.spark.sql.functions._ // 你的示例数据 val df1 = sc.parallelize(List((1, 1), (1, 1), (1, 1), (1, 2),(1, 3), (2, 1), (2, 2), (2, 2))).toDF("key1","key2") // 定义窗口:按key1和key2分组 val keyGroupWindow = Window.partitionBy("key1", "key2") // 添加计数列 val resultDF = df1.withColumn("key2_count", count("*").over(keyGroupWindow)) // 查看结果 resultDF.show()
运行后会得到你期望的输出:
+----+----+----------+ |key1|key2|key2_count| +----+----+----------+ | 1| 1| 3| | 1| 1| 3| | 1| 1| 3| | 1| 2| 1| | 1| 3| 1| | 2| 1| 1| | 2| 2| 2| | 2| 2| 2| +----+----+----------+
这个方法的核心是partitionBy("key1", "key2"),它会把数据分成每个(key1, key2)的子组,然后count("*").over(...)计算每个子组的行数,也就是该组合的出现次数,并将这个值赋给子组内的每一行。
方法二:分组统计后关联原数据
如果你更习惯用分组+join的方式,也可以先统计每个(key1, key2)的计数,再通过join把结果合并回原DataFrame:
import org.apache.spark.sql.functions._ val df1 = sc.parallelize(List((1, 1), (1, 1), (1, 1), (1, 2),(1, 3), (2, 1), (2, 2), (2, 2))).toDF("key1","key2") // 先统计每个(key1, key2)的出现次数 val countDF = df1.groupBy("key1", "key2") .agg(count("*").alias("key2_count")) // 左关联原数据,将计数列添加到每一行 val resultDF = df1.join(countDF, Seq("key1", "key2"), "left") resultDF.show()
这个方法的结果和窗口函数完全一致,适合需要单独保存统计结果的场景。
原代码的问题说明
你原来的代码df2.groupBy("key1").groupBy("key2").agg(count("key2")).show()存在两个问题:
- 连续调用
groupBy并不是同时按两个字段分组,第二次groupBy("key2")是对第一次groupBy("key1")后的聚合结果再次分组,这完全不是你需要的按(key1, key2)组合统计的逻辑; - 即使分组逻辑正确,这种方式得到的是聚合后的统计结果,无法直接关联回原DataFrame的每一行,也就无法实现给每条数据添加对应计数的需求。
内容的提问来源于stack exchange,提问作者Ryan Wang
相关产品推荐
相关产品推荐

