Polars分组上下文使用filter出现不确定结果的原因咨询
Polars分组聚合结果不稳定的原因分析
问题代码
import polars as pl df = pl.DataFrame( { "day": [2, 2, 2, 2, 2, 2, 1, 1], "y": [4, 5, 8, 7, 9, None, None, None], "x": [1, 2, 3, 4, 5, 6, 1, 2], } ) xcol = "x" ycol = "y" f = pl.col(ycol).is_not_null() & pl.col(xcol).is_not_null() df.groupby("day").agg( (pl.col(xcol) - pl.col(xcol).filter(f).mean()).filter(f).sum().alias("filtered_sum") )
出现的两种结果
第一种执行结果:
day filtered_sum 1 null 2 -3.0
第二种执行结果:
day filtered_sum 2 0.0 1 null
期望结果
day filtered_sum 2 0.0 1 null
原因分析
问题核心在于过滤条件f的作用范围未明确绑定到分组内部:
- 你定义的
f是全局列表达式,没有限定它只在分组后的数据上生效。 - Polars的查询优化器会根据内部调度逻辑,可能选择两种不同的执行顺序:
- 先对整个数据集执行
filter(f),计算全局的x均值后再分组求和,此时差值的和会是错误的-3.0; - 先按
day分组,再在每个分组内执行filter(f)并计算组内均值,这时候符合数学逻辑——一组数据减去自身均值的和必然为0,得到正确的0.0。
- 先对整个数据集执行
这种执行顺序的不确定性,导致了多次运行代码出现不同结果。
修复后的代码
要确保过滤和均值计算都在分组内部进行,需要将过滤逻辑明确嵌套到分组的上下文里,比如直接在聚合表达式中写过滤条件:
import polars as pl df = pl.DataFrame( { "day": [2, 2, 2, 2, 2, 2, 1, 1], "y": [4, 5, 8, 7, 9, None, None, None], "x": [1, 2, 3, 4, 5, 6, 1, 2], } ) xcol = "x" ycol = "y" df.groupby("day").agg( (pl.col(xcol) - pl.col(xcol).filter(pl.col(ycol).is_not_null() & pl.col(xcol).is_not_null()).mean()) .filter(pl.col(ycol).is_not_null() & pl.col(xcol).is_not_null()) .sum() .alias("filtered_sum") )
或者用pl.when更简洁地限定处理范围:
df.groupby("day").agg( pl.when(pl.col(ycol).is_not_null() & pl.col(xcol).is_not_null()) .then(pl.col(xcol) - pl.col(xcol).filter(pl.col(ycol).is_not_null() & pl.col(xcol).is_not_null()).mean()) .sum() .alias("filtered_sum") )
修改后无论执行多少次,都会得到你期望的结果。
内容的提问来源于stack exchange,提问作者Keptain
相关产品推荐
相关产品推荐

