如何在Polars中按example_id分组混洗并拆分DataFrame?
Polars实现按示例分组混洗与数据集拆分
核心思路
要实现按example_id分组混洗(保留组内行顺序),并基于完整示例拆分训练/验证/测试集,全程用Polars LazyFrame操作,无需手动转换example_id为数组,可分两步完成:
1. 按示例分组混洗数据
通过生成打乱的example_id排序键,关联回原数据后排序,即可保证组内连续且整体顺序随机:
import polars as pl # 构造示例LazyFrame df = pl.LazyFrame({ "example_id": [1, 1, 2, 2, 2, 3, 3, 3, 4, 4], "other_col": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10] }) # 生成打乱后的example_id映射(带排序键) shuffled_ids = df.select("example_id").unique().with_row_index("shuffle_order").shuffle() # 按打乱后的组顺序重新排列原数据,保留组内行的原始顺序 shuffled_df = df.join(shuffled_ids, on="example_id").sort(["shuffle_order", "example_id"])
2. 按比例拆分完整示例
基于打乱后的example_id列表,按比例分配拆分标签,再关联回原数据完成拆分:
# 计算各数据集的示例数量 total_ids = shuffled_ids.select(pl.count()).collect().item() train_size = int(total_ids * 0.6) val_size = int(total_ids * 0.2) # 为每个示例分配拆分标签 split_ids = shuffled_ids.with_columns( pl.when(pl.col("shuffle_order") < train_size).then("train") .when(pl.col("shuffle_order") < train_size + val_size).then("val") .otherwise("test").alias("split") ) # 关联拆分标签并拆分数据集 split_df = shuffled_df.join(split_ids, on=["example_id", "shuffle_order"]) train_df = split_df.filter(pl.col("split") == "train").drop(["shuffle_order", "split"]) val_df = split_df.filter(pl.col("split") == "val").drop(["shuffle_order", "split"]) test_df = split_df.filter(pl.col("split") == "test").drop(["shuffle_order", "split"])
关键说明
- 全程基于LazyFrame操作,延迟计算保证性能;
- 混洗时通过
shuffle()打乱唯一example_id,再关联排序,既保证组内顺序不变,又实现整体组的随机排列; - 拆分时直接基于打乱后的
example_id列表分配标签,确保每个示例被完整分配到单一数据集。
内容的提问来源于stack exchange,提问作者bkw1491
相关产品推荐
相关产品推荐

