如何用DataFrame按col1分组获取col2计数Top3的结果?
问题:按col1分组获取col2计数Top3的行
我需要基于col2的计数,为col1中的每个分组获取col2的Top3行。
原始表结构及数据
| col1 | col2 |
|---|---|
| A | 1 |
| B | 2 |
| A | 2 |
| B | 2 |
| B | 1 |
| B | 1 |
| B | 1 |
| A | 3 |
| A | 2 |
| B | 4 |
| A | 2 |
| B | 2 |
| A | 3 |
| A | 4 |
例如,分组A中col2=1出现1次,col2=2出现3次,col2=3出现2次(分组B类似)。
期望输出
| col1 | col2 | count(col2) |
|---|---|---|
| A | 2 | 3 |
| A | 3 | 2 |
| A | 1 | 1 |
| B | 1 | 3 |
| B | 2 | 2 |
| B | 4 | 1 |
已实现的SQL查询
SELECT col1, col2, x FROM (SELECT col1, col2, count(col2) AS x, ROW_NUMBER() OVER (PARTITION BY col1 ORDER BY count(col2) DESC) AS rn FROM data GROUP BY col1, col2) tmp WHERE rn <= 3 ORDER BY col1
尝试的DataFrame代码(未得到预期结果)
df.withColumn("rank",dense_rank().over(Window.partitionBy("col1")) .filter(col("rank")<=3) .groupby(col1,col2) .agg(first("col2")) .show()
目前最接近的代码(问题:展示所有行,无法仅保留每组Top3)
df.groupBy("col1","col2").count() .withColumn("rank",rank().over(Window.partitionBy("col1").orderBy(desc("count")))) .where("rank <= 3") .show()
解决方案
问题出在使用rank()函数上:当存在计数相同的行时,rank()会给它们分配相同的排名,导致最终结果中每组的行数可能超过3行。而你的SQL逻辑中使用的是ROW_NUMBER(),它会给每行分配唯一的排名,即使计数相同也会按顺序编号,这样就能严格保留每组Top3。
修改后的正确代码如下:
from pyspark.sql import Window from pyspark.sql.functions import desc, row_number # 先按col1、col2分组统计计数 count_df = df.groupBy("col1", "col2").count() # 定义窗口:按col1分区,按count降序排序,用row_number生成排名 window_spec = Window.partitionBy("col1").orderBy(desc("count")) # 添加排名列,过滤排名<=3的行,最后排序输出 result_df = count_df.withColumn("rank", row_number().over(window_spec)) \ .filter("rank <= 3") \ .orderBy("col1", "rank") \ .drop("rank") # 可选:如果不需要保留rank列可以删除 result_df.show()
这段代码的逻辑和你实现的SQL完全一致:
- 先分组统计每个col1-col2组合的计数
- 按col1分区,对每个分区内的行按count降序排列,用
row_number()生成唯一排名 - 过滤排名<=3的行,最后按col1和rank排序,得到预期的Top3结果
如果需要处理并列排名的场景(比如计数相同的行都要保留,即使超过3个),可以把row_number()换成dense_rank(),但这和你原始SQL的逻辑不同,需要根据实际需求选择。
内容的提问来源于stack exchange,提问作者mhs
相关产品推荐
相关产品推荐

