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

如何在PySpark中基于目标变量和字段对140K小数据集做分层划分训练、验证、测试集?

Hey there! Let's walk through how to do a stratified train/validation/test split in PySpark for your 140K dataset, preserving the distribution of both your target variable and an additional stratification column. This is crucial to ensure your model sees representative samples across all splits, especially with smaller datasets where random splits can easily skew distributions.

Step 1: Setup & Import Required Modules

First, make sure you have your Spark session initialized and the necessary functions imported:

from pyspark.sql import SparkSession
from pyspark.sql.functions import rand, when, col, count
from pyspark.sql.window import Window

# Initialize Spark session (skip if you already have one running)
spark = SparkSession.builder.appName("StratifiedDataSplit").getOrCreate()

Step 2: Define Your Stratification Columns & Split Ratios

Let's assume your DataFrame is named df, your target variable is target_label, and the additional field you want to stratify on is category_col. We'll use a 70/20/10 split for train/validation/test, but you can adjust these ratios as needed.

Step 3: Perform Stratified Split

The key here is to use a window function partitioned by both stratification columns—this ensures we generate random values within each unique (target, category) group, so every group gets split proportionally.

# Define the window to partition by both your target and stratification column
strat_window = Window.partitionBy("target_label", "category_col")

# Add a random value column (seeded for reproducibility) and assign each row to a dataset
df_with_split = df.withColumn(
    "split_rand",
    rand(seed=42).over(strat_window)  # Seed ensures the same split every time
).withColumn(
    "dataset",
    when(col("split_rand") < 0.7, "train")
    .when(col("split_rand") < 0.9, "val")  # 0.7 to 0.9 covers 20% of the data
    .otherwise("test")
)

Step 4: Split Into Separate DataFrames

Now we can filter to get our three distinct datasets, cleaning up the temporary columns we added:

train_df = df_with_split.filter(col("dataset") == "train").drop("split_rand", "dataset")
val_df = df_with_split.filter(col("dataset") == "val").drop("split_rand", "dataset")
test_df = df_with_split.filter(col("dataset") == "test").drop("split_rand", "dataset")

If your dataset has small groups (e.g., a (target, category) pair with only 2-3 samples), proportional splits might result in some groups missing from validation/test sets. To fix this, we can force small groups to go entirely into the training set:

# First calculate the size of each stratification group
df_with_group_size = df.groupBy("target_label", "category_col") \
    .agg(count("*").alias("group_size")) \
    .join(df, on=["target_label", "category_col"])

# Adjust the split logic to keep small groups in training
df_with_split = df_with_group_size.withColumn(
    "split_rand",
    rand(seed=42).over(strat_window)
).withColumn(
    "dataset",
    when(col("group_size") < 5, "train")  # Groups with <5 samples go to train
    .when(col("split_rand") < 0.7, "train")
    .when(col("split_rand") < 0.9, "val")
    .otherwise("test")
)

Step 6: Verify the Split

To make sure the stratification worked, you can compare the distribution of your stratification columns across splits:

# Calculate total distribution of groups
total_dist = df.groupBy("target_label", "category_col").count().withColumnRenamed("count", "total_count")

# Calculate train distribution and compare
train_dist = train_df.groupBy("target_label", "category_col").count()
train_dist.join(total_dist, on=["target_label", "category_col"]) \
    .withColumn("train_ratio", col("count") / col("total_count")) \
    .show()

You should see ratios close to your desired split (e.g., ~0.7 for train) across all groups.

内容的提问来源于stack exchange,提问作者Jay Gondin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 07:58:52