如何对Parquet数据集进行惰性GroupShuffleSplit划分?
按Match分组的训练测试拆分方案(惰性处理大Parquet数据集)
针对你需要按match_id分组拆分训练/测试集(80% match进训练,20%进测试),且要求惰性评估适配大Parquet数据集的需求,以下是几种可行方案:
方案1:Polars(推荐,适配你的当前环境)
Polars原生支持惰性模式(Lazy API),仅在需要时触发计算,完美适配大文件场景。核心思路是先提取所有唯一match_id并拆分,再通过惰性过滤得到训练/测试集:
import polars as pl from sklearn.model_selection import GroupShuffleSplit # 惰性读取Parquet文件(不加载全量数据到内存) lazy_df = pl.scan_parquet("your_dataset.parquet") # 获取所有唯一match_id(仅触发一次小计算,数据量小) unique_matches = lazy_df.select("match_id").unique().collect().to_series().to_list() # 用GroupShuffleSplit拆分match_id gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42) train_idx, test_idx = next(gss.split(unique_matches, groups=unique_matches)) train_match_ids = [unique_matches[i] for i in train_idx] test_match_ids = [unique_matches[i] for i in test_idx] # 惰性过滤生成训练/测试集(此时仍未加载数据) train_lazy = lazy_df.filter(pl.col("match_id").is_in(train_match_ids)) test_lazy = lazy_df.filter(pl.col("match_id").is_in(test_match_ids)) # 按需触发计算,比如写入Parquet train_lazy.sink_parquet("train_data.parquet") test_lazy.sink_parquet("test_data.parquet")
方案2:Dask(分布式场景适配)
Dask虽无法直接对接sklearn的拆分器,但同样可以通过"拆分分组ID+惰性过滤"的思路实现需求:
import dask.dataframe as dd from sklearn.model_selection import GroupShuffleSplit # 读取Parquet为Dask DataFrame(惰性加载) dask_df = dd.read_parquet("your_dataset.parquet") # 获取唯一match_id(分布式计算,仅返回小批量结果) unique_matches = dask_df["match_id"].unique().compute().tolist() # 拆分match_id gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42) train_idx, test_idx = next(gss.split(unique_matches, groups=unique_matches)) train_match_ids = [unique_matches[i] for i in train_idx] test_match_ids = [unique_matches[i] for i in test_idx] # 惰性过滤得到训练/测试集 train_dask = dask_df[dask_df["match_id"].isin(train_match_ids)] test_dask = dask_df[dask_df["match_id"].isin(test_match_ids)] # 写入结果或执行其他计算 train_dask.to_parquet("train_dask.parquet", write_index=False) test_dask.to_parquet("test_dask.parquet", write_index=False)
方案3:PySpark(超大规模分布式场景)
如果数据集达到PB级,PySpark的分布式惰性计算更适配,实现逻辑一致:
from pyspark.sql import SparkSession from sklearn.model_selection import GroupShuffleSplit # 初始化Spark会话 spark = SparkSession.builder.appName("MatchGroupSplit").getOrCreate() # 惰性读取Parquet spark_df = spark.read.parquet("your_dataset.parquet") # 获取唯一match_id unique_matches = [row.match_id for row in spark_df.select("match_id").distinct().collect()] # 拆分match_id gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42) train_idx, test_idx = next(gss.split(unique_matches, groups=unique_matches)) train_match_ids = [unique_matches[i] for i in train_idx] # 过滤生成训练/测试集 train_spark = spark_df.filter(spark_df.match_id.isin(train_match_ids)) test_spark = spark_df.filter(spark_df.match_id.isin(test_match_ids)) # 写入结果 train_spark.write.parquet("train_spark.parquet") test_spark.write.parquet("test_spark.parquet")
核心逻辑说明
所有方案的核心都是先拆分分组标识(match_id),再过滤原数据集:
- 先提取所有唯一的
match_id(数据量远小于全量数据集,计算成本低) - 用sklearn的
GroupShuffleSplit拆分这些ID,保证分组完整性 - 通过惰性过滤得到训练/测试集,全程不加载全量数据到内存
内容的提问来源于stack exchange,提问作者Vivek Joshy
相关产品推荐
相关产品推荐

