Spark DataFrame按ID分组抽取指定百分比记录实现方法
Spark DataFrame 按ID分组按比例抽样最优实现
方案选型结论
- 不要使用全局
sample()方法:该方法为全数据集按比例随机抽取,无法保证每个ID分组达到指定抽样比例,小体量分组甚至可能抽不到任何记录 - 不推荐
groupBy()搭配UDF/Pandas UDF实现:这类方案需要将同组数据全量汇聚到单执行节点处理,存在严重的性能瓶颈,数据倾斜时极易OOM,且序列化/反序列化开销极高 - 最优生产可用方案:使用原生窗口函数按ID分区,组内随机排序后按比例截取样本,全流程受Spark Catalyst优化器原生支持,无额外性能损耗,兼容Spark 2.x及以上所有版本,抽样均匀性、数量准确性均为最优
代码实现(对应15%抽样需求)
以PySpark API为例,Scala/Java API逻辑完全一致:
from pyspark.sql import SparkSession from pyspark.sql.window import Window import pyspark.sql.functions as F spark = SparkSession.builder.appName("grouped_sampling").getOrCreate() # 构造示例数据集:id=1共3条记录,id=2共2条记录 source_df = spark.createDataFrame([ (1, "x"), (1, "y"), (1, "z"), (2, "x"), (2, "y") ], schema=["id", "ip"]) # 配置抽样比例 SAMPLE_RATE = 0.15 # 核心抽样逻辑 id_window = Window.partitionBy("id").orderBy(F.rand(seed=42)) count_window = Window.partitionBy("id") sampled_result = source_df.withColumn("row_idx", F.row_number().over(id_window))\ .withColumn("group_total", F.count("*").over(count_window))\ .filter(F.col("row_idx") <= F.ceil(F.col("group_total") * SAMPLE_RATE))\ .drop("row_idx", "group_total")
实现细节说明
- 随机排序逻辑:
F.rand(seed=42)为组内每条记录分配0~1区间的均匀随机值,保证组内抽样的随机性,固定seed可复现抽样结果 - 样本数量控制:用
ceil对「组总记录数抽样比例」的结果向上取整,保证记录数极少的小分组(比如示例中id=2仅2条记录,215%=0.3)也能抽到至少1条样本;如果需要严格向下取整、允许小分组无样本,可替换为F.floor - 性能表现:窗口分区逻辑与Spark原生shuffle逻辑对齐,无多余计算开销,亿级数据量下运行稳定性远高于UDF类方案
可选简化方案(Spark 3.1+)
如果对每组样本数的精度要求不高(允许概率意义上符合比例,存在±1~2条的误差),可以用内置的sampleBy分层抽样API简化写法:
# 生成所有ID对应的抽样比例映射 all_ids = [r.id for r in source_df.select("id").distinct().collect()] sample_fractions = {_id: SAMPLE_RATE for _id in all_ids} sampled_result_simple = source_df.sampleBy(col="id", fractions=sample_fractions, seed=42)
注意:
sampleBy底层为伯努利抽样,每组返回的样本数为概率期望值,不是精确值,对样本数量准确性要求高的生产场景优先选择窗口函数方案。
内容的提问来源于stack exchange,提问作者Utkarsh Roy
相关产品推荐
相关产品推荐

