Polars懒加载模式下list.to_struct()抛出"expected known type"异常求助
问题:Polars延迟模式下分组聚合返回嵌套数组结构时的类型错误
这是两个Polars技术问题的后续问询:
- 如何在Polars分组上下文返回多统计量为多列?
- 如何在Polars数据框中展平/拆分数组元组并计算列均值?
问题可通过以下示例代码复现:
from functools import partial import polars as pl import statsmodels.api as sm def ols_stats(s, yvar, xvars): df = s.struct.unnest() reg = sm.OLS(df[yvar].to_numpy(), df[xvars].to_numpy(), missing="drop").fit() return pl.Series(values=(reg.params, reg.tvalues), nan_to_null=True) df = pl.DataFrame( { "day": [1, 1, 1, 1, 1, 2, 2, 2, 2, 2], "y": [1, 6, 3, 2, 8, 4, 5, 2, 7, 3], "x1": [1, 8, 2, 3, 5, 2, 1, 2, 7, 3], "x2": [8, 5, 3, 6, 3, 7, 3, 2, 9, 1], } ).lazy() res = df.group_by("day").agg( pl.struct("y", "x1", "x2") .map_elements(partial(ols_stats, yvar="y", xvars=["x1", "x2"])) .alias("params") ) res.with_columns( pl.col("params").list.eval(pl.element().list.explode()).list.to_struct() ).unnest("params").collect()
运行上述延迟模式代码会抛出错误:
PanicException: expected known type
而移除.lazy()和.collect(),以立即模式运行时,代码可正常执行并得到预期结果:
shape: (2, 5) ┌─────┬──────────┬──────────┬──────────┬───────────┐ │ day ┆ field_0 ┆ field_1 ┆ field_2 ┆ field_3 │ │ --- ┆ --- ┆ --- ┆ --- ┆ --- │ │ i64 ┆ f64 ┆ f64 ┆ f64 ┆ f64 │ ╞═════╪══════════╪══════════╪══════════╪═══════════╡ │ 2 ┆ 0.466089 ┆ 0.503127 ┆ 0.916982 ┆ 1.451151 │ │ 1 ┆ 1.008659 ┆ -0.03324 ┆ 3.204266 ┆ -0.124422 │ └─────┴──────────┴──────────┴──────────┴───────────┘
请问该问题的原因是什么?应如何解决?
原因分析
- Polars的**延迟执行模式(Lazy API)**需要提前明确所有操作的输出类型,而
map_elements默认不会自动推断自定义函数的返回类型。在上述代码中,ols_stats返回的是包含数组的Series,但延迟模式无法确定这个返回值的具体结构(比如数组长度、元素类型),导致后续的list.eval和list.to_struct操作因类型未知而触发PanicException。 - 立即模式(Eager API)是逐组实时执行,能动态推断出返回值的类型,因此可以正常运行。
解决方法
核心是给map_elements指定明确的返回类型,让Polars延迟模式能识别输出结构。具体实现:
- 定义
ols_stats的返回类型:返回值是包含两个长度为2的浮点数组的列表,因此类型为pl.List(pl.List(pl.Float64))。 - 在
map_elements中通过return_dtype参数指定该类型。 - 简化展平操作:用
list.flatten()替代list.eval(pl.element().list.explode()),效果一致且更简洁。
修改后的代码:
from functools import partial import polars as pl import statsmodels.api as sm def ols_stats(s, yvar, xvars): df = s.struct.unnest() reg = sm.OLS(df[yvar].to_numpy(), df[xvars].to_numpy(), missing="drop").fit() return pl.Series(values=(reg.params, reg.tvalues), nan_to_null=True) df = pl.DataFrame( { "day": [1, 1, 1, 1, 1, 2, 2, 2, 2, 2], "y": [1, 6, 3, 2, 8, 4, 5, 2, 7, 3], "x1": [1, 8, 2, 3, 5, 2, 1, 2, 7, 3], "x2": [8, 5, 3, 6, 3, 7, 3, 2, 9, 1], } ).lazy() # 指定return_dtype明确返回类型 res = df.group_by("day").agg( pl.struct("y", "x1", "x2") .map_elements( partial(ols_stats, yvar="y", xvars=["x1", "x2"]), return_dtype=pl.List(pl.List(pl.Float64)) ) .alias("params") ) # 简化展平操作 res.with_columns( pl.col("params").list.flatten().list.to_struct() ).unnest("params").collect()
运行修改后的代码,延迟模式下即可得到与立即模式一致的预期结果。
内容的提问来源于stack exchange,提问作者lebesgue
相关产品推荐
相关产品推荐

