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

如何将442亿行DataFrame无重叠划分为训练、验证、测试集?

处理442亿行DataFrame的无重叠划分方案

首先得说,你遇到的核心问题是超大数据量下的内存限制和划分逻辑的严谨性,先拆解下你之前两种方法的问题:

你的两种方法的问题分析

  • 方法一:df.sample(frac=1)会把整个442亿行的DataFrame全量打乱并生成副本,这在单节点环境下完全不可能——内存根本装不下这么大的数据,甚至len(df)计算都会因为行数过多导致性能问题或直接报错。虽然划分逻辑上不会重叠,但执行层面完全不可行。
  • 方法二:三次独立调用np.random.rand(len(df)),相当于给每行生成了三个独立的随机数,这会导致两个严重问题:
    1. 部分行可能同时满足多个划分条件(比如某行第一次随机数<0.8进train,第二次又>0.8且<0.95进valid),造成数据重叠;
    2. 部分行可能三个条件都不满足(比如随机数刚好在某个边界的极端情况),导致数据遗漏。

适合超大数据量的无重叠划分方案

针对442亿行这种规模的数据,单节点Pandas肯定扛不住,必须用分块处理或分布式计算框架,核心思路是:给每行分配唯一的一个随机数,然后根据这个随机数的范围划分到对应的集合,确保每行只属于一个组,同时避免全量加载数据。

方案1:Pandas分块处理(适合小集群/单节点但数据分文件存储的情况)

如果你的数据是拆分成多个文件(比如多个CSV/Parquet),可以逐文件分块读取,给每个块生成一次随机数后划分,再分别保存:

import pandas as pd
import numpy as np
from glob import glob

# 替换为你的数据文件路径
data_files = glob("/path/to/your/data/*.parquet")

# 临时文件列表,用于后续合并
train_temp_files = []
valid_temp_files = []
test_temp_files = []

for file in data_files:
    # 分块读取,chunksize根据你的内存调整,比如100万行/块
    chunk_iter = pd.read_parquet(file, chunksize=1_000_000)
    for chunk in chunk_iter:
        # 给当前块的所有行生成唯一随机数
        chunk["split_tag"] = np.random.rand(len(chunk))
        
        # 按比例划分
        train_chunk = chunk[chunk["split_tag"] < 0.8].drop(columns=["split_tag"])
        valid_chunk = chunk[(chunk["split_tag"] >= 0.8) & (chunk["split_tag"] < 0.95)].drop(columns=["split_tag"])
        test_chunk = chunk[chunk["split_tag"] >= 0.95].drop(columns=["split_tag"])
        
        # 保存临时文件(用随机后缀避免重名)
        train_path = f"/tmp/train_{np.random.randint(100000)}.parquet"
        valid_path = f"/tmp/valid_{np.random.randint(100000)}.parquet"
        test_path = f"/tmp/test_{np.random.randint(100000)}.parquet"
        
        train_chunk.to_parquet(train_path, index=False)
        valid_chunk.to_parquet(valid_path, index=False)
        test_chunk.to_parquet(test_path, index=False)
        
        train_temp_files.append(train_path)
        valid_temp_files.append(valid_path)
        test_temp_files.append(test_path)

# 合并所有临时文件,得到最终的划分数据集
pd.concat([pd.read_parquet(f) for f in train_temp_files]).to_parquet("/path/to/final/train.parquet", index=False)
pd.concat([pd.read_parquet(f) for f in valid_temp_files]).to_parquet("/path/to/final/valid.parquet", index=False)
pd.concat([pd.read_parquet(f) for f in test_temp_files]).to_parquet("/path/to/final/test.parquet", index=False)
  • 注意:优先用Parquet格式,比CSV高效太多,读写速度快、占用空间小,还支持分块操作。

方案2:Dask分布式处理(类Pandas语法,适合中大规模数据)

Dask是Pandas的分布式扩展,能处理远超内存的数据集,语法和Pandas几乎一致,自动帮你处理分块和并行计算:

import dask.dataframe as dd
import numpy as np

# 读取数据,支持多种格式,Parquet最优
df = dd.read_parquet("/path/to/your/data/*.parquet")

# 给每行生成唯一随机数,Dask会分布式计算
df["split_tag"] = dd.random.random(size=len(df), chunksize=1_000_000)

# 按比例划分
train = df[df["split_tag"] < 0.8].drop(columns=["split_tag"])
valid = df[(df["split_tag"] >= 0.8) & (df["split_tag"] < 0.95)].drop(columns=["split_tag"])
test = df[df["split_tag"] >= 0.95].drop(columns=["split_tag"])

# 保存划分后的数据集,Dask会自动分块存储
train.to_parquet("/path/to/final/train/", write_index=False)
valid.to_parquet("/path/to/final/valid/", write_index=False)
test.to_parquet("/path/to/final/test/", write_index=False)

方案3:PySpark(适合超大规模数据,442亿行首选)

如果数据量达到百亿级别,PySpark是业界标准的解决方案,依托Spark的分布式集群,轻松处理PB级数据:

from pyspark.sql import SparkSession
from pyspark.sql.functions import rand

# 初始化Spark会话
spark = SparkSession.builder \
    .appName("HugeDataSplit") \
    .getOrCreate()

# 读取数据,支持Parquet、CSV等格式,Parquet优先
df = spark.read.parquet("/path/to/your/data/")

# 添加随机数列,rand()会给每行生成0-1之间的唯一随机数
df = df.withColumn("split_tag", rand())

# 按比例划分数据集
train_df = df.filter(df.split_tag < 0.8).drop("split_tag")
valid_df = df.filter((df.split_tag >= 0.8) & (df.split_tag < 0.95)).drop("split_tag")
test_df = df.filter(df.split_tag >= 0.95).drop("split_tag")

# 保存数据,用Parquet格式,mode="overwrite"表示覆盖已有文件
train_df.write.parquet("/path/to/final/train/", mode="overwrite")
valid_df.write.parquet("/path/to/final/valid/", mode="overwrite")
test_df.write.parquet("/path/to/final/test/", mode="overwrite")

# 关闭Spark会话
spark.stop()

验证划分结果的正确性

划分完成后,一定要验证是否存在重叠或遗漏,以PySpark为例:

# 检查train和valid是否有重叠的id
train_df.select("id").intersect(valid_df.select("id")).count()  # 结果应为0
# 检查train和test是否有重叠的id
train_df.select("id").intersect(test_df.select("id")).count()   # 结果应为0
# 检查valid和test是否有重叠的id
valid_df.select("id").intersect(test_df.select("id")).count()   # 结果应为0
# 检查总数据量是否和原数据一致
total_rows = train_df.count() + valid_df.count() + test_df.count()
original_rows = df.count()
print(total_rows == original_rows)  # 结果应为True

内容的提问来源于stack exchange,提问作者John Davis

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:11:28