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

基于code列的PySpark数据集Train/Eval/Test划分异常问题

解决Spark DataFrame按code划分Train/Eval/Test时code跨集的问题

问题背景

需要基于Spark DataFrame的code列完成Train/Eval/Test数据集划分,要求同一code的所有行必须归属于同一集合,不得跨集。现有实现支持手动指定部分code归入对应集合,但验证发现部分code同时出现在多个集合中。

问题根源

  1. 手动指定的code列表存在交集:如果同一个code同时出现在codes_to_put_in_train、codes_to_put_in_eval或codes_to_put_in_test中的多个列表里,会导致该code被重复匹配。
  2. 随机分配逻辑未排除已指定的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 17:27:07