PySpark如何实现sklearn中StratifiedGroupKFold的同等功能?
PySpark实现分层分组隔离的训练测试集拆分方案
核心思路:必须先在分组维度做分层采样,再关联回原始样本表,即可同时满足「同组不跨拆分集」和「标签分布分层」两个要求,完全对齐StratifiedGroupKFold的拆分逻辑。
具体实现步骤
假设你的原始数据集为df,分组标识列为group_id,二分类标签列为label,按7:3拆分训练集/测试集,实现代码如下:
- 先聚合得到每个分组的分层键
分层键通常取分组内占比最高的标签,若同一组内所有样本标签完全一致,可直接用
F.first("label")替代F.mode("label")提升计算效率
from pyspark.sql import functions as F group_label_df = df.groupBy("group_id").agg( F.mode("label").alias("group_stratify_key") )
- 对分组级表做分层采样,得到训练集对应的分组ID集合
train_split_ratio = 0.7 train_groups = group_label_df.stat.sampleBy( col="group_stratify_key", # 二分类两个类别均按训练占比采样,保证整体标签分布与原数据集一致 fractions={0: train_split_ratio, 1: train_split_ratio}, seed=42 # 固定随机种子保证拆分结果可复现 ).select("group_id")
- 关联回原始表得到最终训练、测试集
# 训练集:包含所有训练分组下的样本 train_df = df.join(train_groups, on="group_id", how="inner") # 测试集:包含所有非训练分组下的样本 test_df = df.join(train_groups, on="group_id", how="left_anti")
扩展说明
- 如果需要实现K折分层分组交叉验证,仅需修改第二步逻辑:给每个分组按分层键随机分配1~K的折数编号,后续按编号筛选对应折的训练、测试集即可
- 若存在多分类场景,仅需调整
sampleBy的fractions参数,给每个类别配置相同的采样比例即可
内容的提问来源于stack exchange,提问作者michen00
相关产品推荐
相关产品推荐

