如何在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")
Step 5: Handle Small Groups (Optional but Recommended)
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

