Python Polars中使用.over()按指定列分组累加首元素的问题
Python Polars 分组累加筛选完整分组解决方案
问题分析
你原来的代码错误在于:pl.col("pty_nber").first().over("Declaration")会将每个分组的pty_nber首元素广播到该组的所有行,随后的cum_sum()是对整个数据集的所有行进行累加,导致同一个分组的首元素被重复累加多次(组内有多少行就加多少次),这才出现了4+4+4+4+7…的错误结果。
正确实现步骤
要实现按分组首元素累加、筛选累加和低于30的完整分组,需要分两步处理:先计算分组首元素的累加和并筛选有效分组,再用有效分组过滤原数据集:
- 提取每个分组的
Declaration和对应的pty_nber首元素,保留分组顺序:
group_first = dataset.group_by("Declaration", maintain_order=True).agg( pl.col("pty_nber").first().alias("first_pty") )
- 对分组首元素计算累加和:
group_first = group_first.with_columns( cum_sum=pl.col("first_pty").cum_sum() )
- 筛选出累加和小于30的分组:
valid_groups = group_first.filter(pl.col("cum_sum") < 30)["Declaration"]
- 用有效分组过滤原数据集,保留完整分组的所有行:
result = dataset.filter(pl.col("Declaration").is_in(valid_groups))
合并简化写法
也可以把步骤合并成链式调用,更简洁:
result = dataset.filter( pl.col("Declaration").is_in( dataset.group_by("Declaration", maintain_order=True) .agg(pl.col("pty_nber").first()) .with_columns(cum_sum=pl.col("pty_nber").cum_sum()) .filter(pl.col("cum_sum") < 30) .select("Declaration") ) )
内容的提问来源于stack exchange,提问作者McNickSisto
相关产品推荐
相关产品推荐

