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

PySpark如何实现sklearn中StratifiedGroupKFold的同等功能?

PySpark实现分层分组隔离的训练测试集拆分方案

核心思路:必须先在分组维度做分层采样,再关联回原始样本表,即可同时满足「同组不跨拆分集」和「标签分布分层」两个要求,完全对齐StratifiedGroupKFold的拆分逻辑。

具体实现步骤

假设你的原始数据集为df,分组标识列为group_id,二分类标签列为label,按7:3拆分训练集/测试集,实现代码如下:

  1. 先聚合得到每个分组的分层键

分层键通常取分组内占比最高的标签,若同一组内所有样本标签完全一致,可直接用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")
)
  1. 对分组级表做分层采样,得到训练集对应的分组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")
  1. 关联回原始表得到最终训练、测试集
# 训练集:包含所有训练分组下的样本
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 12:24:03