You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用DataFrame按col1分组获取col2计数Top3的结果?

问题:按col1分组获取col2计数Top3的行

我需要基于col2的计数,为col1中的每个分组获取col2的Top3行。

原始表结构及数据

col1col2
A1
B2
A2
B2
B1
B1
B1
A3
A2
B4
A2
B2
A3
A4

例如,分组A中col2=1出现1次,col2=2出现3次,col2=3出现2次(分组B类似)。

期望输出

col1col2count(col2)
A23
A32
A11
B13
B22
B41

已实现的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完全一致:

  1. 先分组统计每个col1-col2组合的计数
  2. 按col1分区,对每个分区内的行按count降序排列,用row_number()生成唯一排名
  3. 过滤排名<=3的行,最后按col1和rank排序,得到预期的Top3结果

如果需要处理并列排名的场景(比如计数相同的行都要保留,即使超过3个),可以把row_number()换成dense_rank(),但这和你原始SQL的逻辑不同,需要根据实际需求选择。

内容的提问来源于stack exchange,提问作者mhs

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.09 03:46:12