Polars升级后.list.sum()返回异常结果,是否存在问题?
Polars 26→29升级后
.list.sum()返回类型异常问题解析 问题现象
将Polars从26版本升级至29版本后,在pl.when().then()分支中使用.list.sum()时出现返回类型异常:
- 在
when/then逻辑内,pl.concat_list(pl.col('A', 'B')).list.sum()返回list[f64]类型(如[5.0]),旧版本则返回预期的f64标量; - 直接在
with_columns中使用相同的.list.sum()逻辑,却能正常返回f64标量。
复现代码:
import polars as pl x = pl.DataFrame( { 'A': [1., None, 3.], 'B': [4., 5., 6.], 'C': [7., 8., None], } ) x.with_columns( pl.when(pl.sum_horizontal('B', 'C') > 12) .then(pl.sum_horizontal(pl.col('A', 'B'))) .otherwise(None) .alias('A+B when B+C>12 (expected)'), pl.when(pl.sum_horizontal('B', 'C') > 12) .then(pl.concat_list(pl.col('A', 'B')).list.sum()) .otherwise(None) .alias('A+B when B+C>12 (odd)'), pl.concat_list(pl.exclude('C')).list.sum().alias('A+B'), )
运行结果:
┌──────┬─────┬──────┬────────────────────────────┬───────────────────────┬─────┐ │ A ┆ B ┆ C ┆ A+B when B+C>12 (expected) ┆ A+B when B+C>12 (odd) ┆ A+B │ │ --- ┆ --- ┆ --- ┆ --- ┆ --- ┆ --- │ │ f64 ┆ f64 ┆ f64 ┆ f64 ┆ list[f64] ┆ f64 │ ╞══════╪═════╪══════╪════════════════════════════╪═══════════════════════╪═════╡ │ 1.0 ┆ 4.0 ┆ 7.0 ┆ null ┆ null ┆ 5.0 │ │ null ┆ 5.0 ┆ 8.0 ┆ 5.0 ┆ [5.0] ┆ 5.0 │ │ 3.0 ┆ 6.0 ┆ null ┆ null ┆ null ┆ 9.0 │ └──────┴─────┴──────┴────────────────────────────┴───────────────────────┴─────┘
原因分析
这是Polars 29版本中上下文类型推断逻辑的变更/潜在bug:
当when/then分支存在otherwise(None)时,Polars的类型推断会尝试兼容then分支结果与None的类型。此处错误地将.list.sum()返回的标量值包装为列表,以适配“空列表+None”的兼容场景,但实际上.list.sum()的预期返回值是标量,而非列表。
而直接在with_columns中使用.list.sum()时,没有None分支的干扰,类型推断能正确识别出结果为标量,因此返回正常的数值类型。
解决方案
方案1:强制指定返回类型
在then分支的.list.sum()后添加.cast(pl.Float64),明确告诉Polars返回标量类型:
pl.when(pl.sum_horizontal('B', 'C') > 12) .then(pl.concat_list(pl.col('A', 'B')).list.sum().cast(pl.Float64)) .otherwise(None) .alias('A+B when B+C>12 (fixed)'),
方案2:改用更直接的sum_horizontal
如果只是需要横向求和,完全可以用sum_horizontal替代concat_list().list.sum(),从根源避免类型推断问题:
pl.when(pl.sum_horizontal('B', 'C') > 12) .then(pl.sum_horizontal(pl.col('A', 'B'))) .otherwise(None) .alias('A+B when B+C>12 (expected)'),
方案3:提取列表中的标量(兜底写法)
若必须使用列表操作,可在.list.sum()后添加.list.first()提取唯一的标量值:
pl.when(pl.sum_horizontal('B', 'C') > 12) .then(pl.concat_list(pl.col('A', 'B')).list.sum().list.first()) .otherwise(None) .alias('A+B when B+C>12 (fixed)'),
关于.list.sum()的预期行为
.list.sum()作为列表的聚合函数,理应始终返回单个标量值,而非列表。本次出现的异常属于Polars版本升级中的类型推断bug,你可以在Polars的官方GitHub仓库提交issue反馈该问题。
内容的提问来源于stack exchange,提问作者Xiaofeng Lu
相关产品推荐
相关产品推荐

