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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 05:00:02