Polars分组取值问题:解决n大于组内元素数的越界错误
在Polars中处理分组gather时的索引越界问题
当在Polars中按分组提取元素时,如果指定的索引超出了分组内的元素总数,会触发OutOfBoundsError。例如以下代码,分组x=2仅包含1个元素,使用索引2会直接报错:
import polars as pl df = pl.DataFrame(dict(x=[1,1,1,2,3,3,3], y=[1,2,3,4,5,6,7])) df.group_by("x").agg(pl.all().gather([0, 2]))
方法1:过滤超出范围的索引(仅保留有效元素)
通过list.eval在每个分组内部动态筛选出小于组长度的索引,再执行gather操作:
import polars as pl df = pl.DataFrame(dict(x=[1,1,1,2,3,3,3], y=[1,2,3,4,5,6,7])) target_indices = [0, 2] result = df.group_by("x").agg( pl.all().list.eval( pl.element().gather(pl.Series(target_indices).filter(pl.element() < pl.len())) ) ) print(result)
输出结果:
shape: (3, 2) ┌──────┬──────────┐ │ x ┆ y │ │ --- ┆ --- │ │ i64 ┆ list[i64]│ ╞══════╪══════════╡ │ 1 ┆ [1, 3] │ ├╌╌╌╌╌╌┼╌╌╌╌╌╌╌╌╌╌┤ │ 2 ┆ [4] │ ├╌╌╌╌╌╌┼╌╌╌╌╌╌╌╌╌╌┤ │ 3 ┆ [5, 7] │ └──────┴──────────┘
方法2:保留索引位置,缺失补Null
如果需要保留指定的索引位置,对超出范围的索引返回Null,可以使用list.get方法:
import polars as pl df = pl.DataFrame(dict(x=[1,1,1,2,3,3,3], y=[1,2,3,4,5,6,7])) target_indices = [0, 2] result = df.group_by("x").agg( pl.all().list.eval( pl.struct([pl.element().get(idx) for idx in target_indices]) ).list.explode() ).unnest("y") print(result)
输出结果:
shape: (6, 3) ┌──────┬────────┬────────┐ │ x ┆ field_0┆ field_1│ │ --- ┆ --- ┆ --- │ │ i64 ┆ i64 ┆ i64 │ ╞══════╪════════╪════════╡ │ 1 ┆ 1 ┆ 3 │ ├╌╌╌╌╌╌┼╌╌╌╌╌╌╌╌┼╌╌╌╌╌╌╌╌┤ │ 2 ┆ 4 ┆ null │ ├╌╌╌╌╌╌┼╌╌╌╌╌╌╌╌┼╌╌╌╌╌╌╌╌┤ │ 3 ┆ 5 ┆ 7 │ └──────┴────────┴────────┘
(注:如果不需要重复行,可将pl.all()替换为pl.all().first()来避免分组内的重复计算)
内容的提问来源于stack exchange,提问作者spitfiredd
相关产品推荐
相关产品推荐

