如何将442亿行DataFrame无重叠划分为训练、验证、测试集?
处理442亿行DataFrame的无重叠划分方案
首先得说,你遇到的核心问题是超大数据量下的内存限制和划分逻辑的严谨性,先拆解下你之前两种方法的问题:
你的两种方法的问题分析
- 方法一:
df.sample(frac=1)会把整个442亿行的DataFrame全量打乱并生成副本,这在单节点环境下完全不可能——内存根本装不下这么大的数据,甚至len(df)计算都会因为行数过多导致性能问题或直接报错。虽然划分逻辑上不会重叠,但执行层面完全不可行。 - 方法二:三次独立调用
np.random.rand(len(df)),相当于给每行生成了三个独立的随机数,这会导致两个严重问题:- 部分行可能同时满足多个划分条件(比如某行第一次随机数<0.8进train,第二次又>0.8且<0.95进valid),造成数据重叠;
- 部分行可能三个条件都不满足(比如随机数刚好在某个边界的极端情况),导致数据遗漏。
适合超大数据量的无重叠划分方案
针对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
相关产品推荐
相关产品推荐

