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

如何对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),再过滤原数据集:

  1. 先提取所有唯一的match_id(数据量远小于全量数据集,计算成本低)
  2. 用sklearn的GroupShuffleSplit拆分这些ID,保证分组完整性
  3. 通过惰性过滤得到训练/测试集,全程不加载全量数据到内存

内容的提问来源于stack exchange,提问作者Vivek Joshy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 01:02:06