Spark中如何组合使用groupBy与sampleBy实现分组内分层采样
实现方案说明
Spark原生GroupedData对象没有直接提供sampleBy方法,你可以通过以下两种常用方案实现「按prod_name分组后,对每个分组内的colour列执行分层采样」的需求:
方案1:窗口函数实现(适合全量大数据量场景,性能更好)
该方案通过窗口标记分组内的随机排序序号,按预设比例过滤行,是精确采样,性能更优:
from pyspark.sql import SparkSession from pyspark.sql.functions import rand, row_number, coalesce, lit from pyspark.sql.window import Window # 初始化Spark会话 spark = SparkSession.builder.appName("grouped_stratified_sample").getOrCreate() # 构造示例数据集 data = [ ("A", "blue", 100, "Y"), ("A", "blue", 200, "N"), ("A", "blue", 300, "Y"), ("A", "blue", 400, "Y"), ("A", "yellow", 500, "N"), ("B", "green", 600, "Y"), ("B", "green", 650, "Y"), ("B", "blue", 700, "N"), ("C", "red", 800, "Y"), ("C", "blue", 900, "N"), ("C", "green", 1000, "N") ] df = spark.createDataFrame(data, schema=["prod_name", "colour", "value", "code"]) # 定义分层采样比例,未出现的颜色默认采样比例为0 frac_map = {"blue":0.5, "yellow":0.1, "green":0.3} default_frac = 0.0 # 按prod_name+colour分组,给组内每行打随机序号、统计组内总行数 window_spec = Window.partitionBy("prod_name", "colour").orderBy(rand()) df = df.withColumn("rn", row_number().over(window_spec))\ .withColumn("group_total", row_number().over(window_spec.rangeBetween(Window.unboundedPreceding, Window.unboundedFollowing))) # 按比例过滤符合要求的行 result = df.filter(df.rn <= df.group_total * coalesce(df.colour.map(frac_map), lit(default_frac)))\ .drop("rn", "group_total") result.show()
方案2:groupBy + applyInPandas(适合逻辑灵活的场景,代码更易读)
该方案对每个prod_name分组转pandas DataFrame后单独采样,适合需要自定义复杂采样逻辑的场景:
import pandas as pd # 定义采样规则 frac_map = {"blue":0.5, "yellow":0.1, "green":0.3} default_frac = 0.0 # 定义分组采样函数 def stratified_sample_per_group(pdf: pd.DataFrame) -> pd.DataFrame: return pdf.groupby("colour", group_keys=False).apply( lambda x: x.sample(frac=frac_map.get(x.name, default_frac), random_state=42) ) # 定义输出schema output_schema = "prod_name string, colour string, value int, code string" # 分组执行采样 result = df.groupby("prod_name").applyInPandas(stratified_sample_per_group, schema=output_schema) result.show()
注意事项
- 采样比例字典需要覆盖所有你需要采样的colour取值,不需要采样的colour可以设置比例为0
- 如果不同prod_name分组需要使用不同的colour采样比例,可将
frac_map改为嵌套字典,按prod_name读取对应比例即可 - 原生
sampleBy是近似采样,如果你需要严格符合比例的采样结果,优先使用方案1
内容的提问来源于stack exchange,提问作者user3735871
相关产品推荐
相关产品推荐

