如何实现Polars LazyFrame的无评估式按比例拆分训练测试集
用Polars LazyFrame实现train_test_split的高效方案
核心思路
给LazyFrame添加一个矢量化生成的随机列,通过计算该列的分位数作为拆分阈值,再用filter分别提取训练集和测试集——全程除了获取阈值的一次轻量操作外,其余均保持懒执行状态,完美匹配需求。
解决代码
import polars as pl def lazy_train_test_split(df: pl.LazyFrame, train_ratio: float = 0.8) -> tuple[pl.LazyFrame, pl.LazyFrame]: # 给LazyFrame添加全局唯一的随机列(矢量化生成,效率拉满) df_rand = df.with_columns(pl.rand().alias("_split_rand")) # 计算拆分阈值:随机列的train_ratio分位数(仅需一次轻量collect,无需全量加载数据) split_threshold = df_rand.select(pl.col("_split_rand").quantile(train_ratio)).collect().item() # 拆分并移除临时随机列 train_df = df_rand.filter(pl.col("_split_rand") <= split_threshold).drop("_split_rand") test_df = df_rand.filter(pl.col("_split_rand") > split_threshold).drop("_split_rand") return train_df, test_df
针对你遇到的问题的分析
df.sample()的局限:确实只能单独抽取训练集,无法直接通过Lazy方式获取剩余的测试集,因为sample是随机抽样而非全量打乱后的拆分。- 添加随机列的效率问题:你之前用逐行映射导致效率低,而Polars内置的
pl.rand()是矢量化操作,底层用C实现,生成随机列的速度和数据量线性相关,完全不会有性能瓶颈。 - 总行数统计的开销:如果数据源是Parquet这类支持元数据的格式,Polars可以直接从元数据获取行数,无需扫描全量数据;如果是CSV等无元数据的格式,统计行数需要一次全量扫描,但开销远小于加载整个数据集。不过上面的方案不需要统计总行数,更适合大数据场景。
测试示例(大数据集验证)
# 生成1000万行的测试LazyFrame big_test_df = pl.LazyFrame({ "feature1": pl.arange(0, 10_000_000), "feature2": pl.randn(10_000_000), "label": pl.randint(0, 2, 10_000_000) }) # 拆分训练集(70%)和测试集(30%) train, test = lazy_train_test_split(big_test_df, train_ratio=0.7) # 验证拆分比例(仅测试时collect,实际使用可保持懒状态) print(f"训练集行数: {train.select(pl.count()).collect().item()}") print(f"测试集行数: {test.select(pl.count()).collect().item()}")
备选方案(基于总行数拆分)
如果你的数据源支持快速获取行数,也可以用全量打乱后按行数拆分,但这种方法需要先统计总行数:
def lazy_train_test_split_by_count(df: pl.LazyFrame, train_ratio: float = 0.8) -> tuple[pl.LazyFrame, pl.LazyFrame]: total_rows = df.select(pl.count()).collect().item() train_size = int(total_rows * train_ratio) # 全量打乱后取前train_size行作为训练集,剩余作为测试集 shuffled = df.sample(fraction=1.0, shuffle=True) return shuffled.head(train_size), shuffled.tail(total_rows - train_size)
但这种方法的shuffle操作开销比随机列拆分更大,因为需要生成全量数据的随机排列,而随机列拆分仅需过滤,更推荐前者。
内容的提问来源于stack exchange,提问作者TomNorway
相关产品推荐
相关产品推荐

