PySpark按item_id分组后保留每组得分前50%记录的实现
PySpark按分组百分比筛选高得分记录的实现方案
需求说明
针对给定的PySpark DataFrame,需按item_id分组,剔除每组中得分最低的50%记录,保留得分最高的50%,同时确保所有item_id都存在(哪怕组内仅1条记录也需保留)。
实现思路
通过两个窗口函数配合完成动态筛选:
- 第一个窗口按
item_id分组,统计每组的总记录数,用于计算动态保留阈值; - 第二个窗口按
item_id分组、score降序排序,给每条记录生成排名; - 计算每组需保留的记录数(总条数的50%,奇数条向上取整),过滤出排名在阈值内的记录。
完整代码
from pyspark.sql import Window from pyspark.sql.functions import col, row_number, count, ceil # 定义窗口:仅按item_id分组,用于统计每组总记录数 window_count = Window.partitionBy("item_id") # 定义窗口:按item_id分组、score降序排序,用于生成排名 window_rank = Window.partitionBy("item_id").orderBy(col("score").desc()) # 数据处理流程 result_df = df.withColumn("total_records", count("*").over(window_count)) \ .withColumn("rank", row_number().over(window_rank)) \ .withColumn("keep_threshold", ceil(col("total_records") * 0.5)) \ .filter(col("rank") <= col("keep_threshold")) \ .drop("total_records", "rank", "keep_threshold")
代码细节说明
count("*").over(window_count):动态计算每个item_id组的总记录数,替代固定值筛选,适配不同组的规模差异;row_number().over(window_rank):按得分从高到低为每组记录排名,确保优先保留高得分数据;ceil(col("total_records") * 0.5):对每组保留数量向上取整,避免因奇数条记录导致item_id被完全剔除(比如3条记录保留2条,1条记录保留1条);- 最后过滤掉超阈值记录,并删除中间计算列,得到符合要求的结果。
超大规模数据集优化提示
针对80亿条的数据集,建议:
- 调整
spark.sql.shuffle.partitions参数(默认200),根据集群资源合理设置,降低shuffle压力; - 提前对
item_id做分区预处理,避免数据倾斜影响计算效率。
内容的提问来源于stack exchange,提问作者Chris_007
相关产品推荐
相关产品推荐

