基于code列的PySpark数据集Train/Eval/Test划分异常问题
解决Spark DataFrame按code划分Train/Eval/Test时code跨集的问题
问题背景
需要基于Spark DataFrame的code列完成Train/Eval/Test数据集划分,要求同一code的所有行必须归属于同一集合,不得跨集。现有实现支持手动指定部分code归入对应集合,但验证发现部分code同时出现在多个集合中。
问题根源
- 手动指定的code列表存在交集:如果同一个
code同时出现在codes_to_put_in_train、codes_to_put_in_eval或codes_to_put_in_test中的多个列表里,会导致该code被重复匹配。 - 随机分配逻辑未排除已指定的code:原代码中,手动指定的code在匹配完第一个
when后,后续的随机判断分支未被跳过,若rand_val刚好落在其他集合的阈值区间,可能被重复标记类型。
解决办法
1. 先校验手动指定的code列表,确保两两无交集
在执行划分前,先检查三个手动列表是否有重叠元素,从源头避免错误:
# 检查手动指定列表的交集情况 train_eval_intersect = set(codes_to_put_in_train) & set(codes_to_put_in_eval) train_test_intersect = set(codes_to_put_in_train) & set(codes_to_put_in_test) eval_test_intersect = set(codes_to_put_in_eval) & set(codes_to_put_in_test) if train_eval_intersect: raise ValueError(f"code同时出现在train和eval列表:{train_eval_intersect}") if train_test_intersect: raise ValueError(f"code同时出现在train和test列表:{train_test_intersect}") if eval_test_intersect: raise ValueError(f"code同时出现在eval和test列表:{eval_test_intersect}")
2. 修改data_type赋值逻辑,隔离手动指定与随机分配的code
在随机分配分支中,仅对未被手动指定的code进行判断,确保每个code只被分配一次类型:
from pyspark.sql import functions as F def train_eval_test_split(df, EVAL_FRACTION=0.15, TEST_FRACTION=0.15): # 先校验手动列表交集 train_eval_intersect = set(codes_to_put_in_train) & set(codes_to_put_in_eval) train_test_intersect = set(codes_to_put_in_train) & set(codes_to_put_in_test) eval_test_intersect = set(codes_to_put_in_eval) & set(codes_to_put_in_test) if train_eval_intersect: raise ValueError(f"code同时出现在train和eval列表:{train_eval_intersect}") if train_test_intersect: raise ValueError(f"code同时出现在train和test列表:{train_test_intersect}") if eval_test_intersect: raise ValueError(f"code同时出现在eval和test列表:{eval_test_intersect}") # 合并所有手动指定的code,用于后续过滤 manual_codes = codes_to_put_in_train + codes_to_put_in_eval + codes_to_put_in_test train_test_split = ( df.select("code") .distinct() .withColumn("rand_val", F.rand(seed=42)) .withColumn( "data_type", F.when(F.col("code").isin(codes_to_put_in_train), "train") .when(F.col("code").isin(codes_to_put_in_eval), "eval") .when(F.col("code").isin(codes_to_put_in_test), "test") # 仅对未手动指定的code执行随机分配 .when( ~F.col("code").isin(manual_codes) & (F.col("rand_val") < TEST_FRACTION), "test" ) .when( ~F.col("code").isin(manual_codes) & (F.col("rand_val") >= TEST_FRACTION) & (F.col("rand_val") < TEST_FRACTION + EVAL_FRACTION), "eval" ) .otherwise("train"), ) ) train_df = train_test_split.filter(F.col("data_type") == "train").join(df, on="code") test_df = train_test_split.filter(F.col("data_type") == "test").join(df, on="code") eval_df = train_test_split.filter(F.col("data_type") == "eval").join(df, on="code") return train_df, eval_df, test_df
关键修改说明
- 新增手动列表交集校验,提前拦截同一code被多次指定的情况。
- 随机分配分支添加
~F.col("code").isin(manual_codes)条件,确保只有未被手动指定的code才进入随机分配逻辑,彻底避免重复标记。
内容的提问来源于stack exchange,提问作者Unguided8018
相关产品推荐
相关产品推荐

