如何高效保留DataFrame重复组的最后3条数据?
高效保留分组后最后N条数据的Spark优化方案
我有如下Spark DataFrame,其中id和country列始终唯一,但存在基于first_name、last_name和sex的重复分组。需要找出这些分组,仅保留每组的最后3条数据,其余删除。
原始DataFrame
| id | first_name | last_name | sex | country |
|---|---|---|---|---|
| 01 | John | Doe | Male | USA |
| 02 | John | Doe | Male | Canada |
| 03 | John | Doe | Male | Mexico |
| 04 | Mark | Kay | Male | Italy |
| 05 | John | Doe | Male | Spain |
| 06 | Mark | Kay | Male | France |
| 07 | John | Doe | Male | Peru |
| 08 | Mark | Kay | Male | India |
| 09 | Mark | Kay | Male | Laos |
| 10 | John | Doe | Male | Benin |
预期处理结果
| id | first_name | last_name | sex | country |
|---|---|---|---|---|
| 05 | John | Doe | Male | Spain |
| 06 | Mark | Kay | Male | France |
| 07 | John | Doe | Male | Peru |
| 08 | Mark | Kay | Male | India |
| 09 | Mark | Kay | Male | Laos |
| 10 | John | Doe | Male | Benin |
当前实现代码
我已经实现了以下代码,能得到预期结果,但数据集规模较大(行数和列数更多),想知道更高效的实现方式:
window_spec = Window.partitionBy('first_name', 'last_name', 'sex').orderBy(F.desc('id')) df_with_row_number = df.withColumn('row_number', F.row_number().over(window_spec)) filtered_df = df_with_row_number.filter('row_number <= 3') result_df = filtered_df.drop('row_number')
优化方案
针对大数据量场景,有以下几种实用优化方向:
1. 减少窗口操作的列数
如果原始DataFrame包含大量非必要列(比如大文本、数组类型字段),可以先提取分组、排序所需的核心字段+主键,筛选完成后再关联回原始数据获取完整列,大幅降低窗口操作的数据处理量:
# 仅保留分组、排序键和唯一主键id temp_df = df.select('id', 'first_name', 'last_name', 'sex') window_spec = Window.partitionBy('first_name', 'last_name', 'sex').orderBy(F.desc('id')) temp_df = temp_df.withColumn('row_num', F.row_number().over(window_spec))\ .filter('row_num <= 3')\ .drop('row_num') # 关联回原始表获取所有字段 result_df = temp_df.join(df, on='id', how='inner')
2. 优化分区并行度
按分组键重分区,让同属一个分组的数据集中在同一个分区内,减少shuffle开销,同时匹配集群资源调整分区数量:
# 按分组键重分区,合理设置分区数(根据集群CPU核数调整) df = df.repartition(200, 'first_name', 'last_name', 'sex') window_spec = Window.partitionBy('first_name', 'last_name', 'sex').orderBy(F.desc('id')) df_with_row_number = df.withColumn('row_number', F.row_number().over(window_spec)) filtered_df = df_with_row_number.filter('row_number <= 3').drop('row_number')
3. 用rank()替代row_number()(无重复排序键时等效)
如果id是严格唯一的,rank()和row_number()输出结果完全一致,部分Spark版本中rank()的底层执行逻辑会略减少排序开销(差异极小,但可尝试):
window_spec = Window.partitionBy('first_name', 'last_name', 'sex').orderBy(F.desc('id')) df_with_rank = df.withColumn('rank', F.rank().over(window_spec)) filtered_df = df_with_rank.filter('rank <= 3').drop('rank')
总结
你的原始实现是Spark处理这类分组筛选需求的标准写法,大数据量下最有效的优化手段是减少窗口操作的列数和合理调整分区策略,这两种方式能直接降低内存占用和shuffle开销,显著提升执行效率。
内容的提问来源于stack exchange,提问作者Kruttika Swaminathan
相关产品推荐
相关产品推荐

