Polars如何解决‘聚合中不允许窗口表达式’报错问题?
在Polars中按分组规则执行回归拟合的问题与解决
问题场景
我需要在Polars DataFrame中,基于specs字典定义的转换规则与自变量列,对转换后的列执行回归分析。以下是简化示例代码:
from functools import partial import polars as pl import numpy as np def ols_fitted(s: pl.Series, yvar: str, xvars: list[str]) -> pl.Series: df = s.struct.unnest() y = df[yvar].to_numpy() X = df[xvars].to_numpy() fitted = np.dot(X, np.linalg.lstsq(X, y, rcond=None)[0]) return pl.Series(values=fitted, nan_to_null=True) df = pl.DataFrame( { "date": [1, 1, 1, 1, 2, 2, 2, 2, 2, 2], "id": [1, 1, 1, 2, 2, 2, 2, 3, 3, 3], "y": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10], "g1": [1, 1, 1, 1, 1, 2, 2, 2, 2, 2], "g2": [1, 1, 1, 2, 2, 2, 3, 3, 3, 3], "g3": [1, 1, 2, 2, 2, 3, 3, 4, 4, 4], "x1": [2, 5, 4, 7, 3, 2, 5, 6, 7, 2], "x2": [1, 5, 3, 4, 5, 6, 4, 3, 2, 1], "x3": [3, 6, 8, 6, 4, 7, 5, 4, 8, 1], } ) specs = { "first": {"yvar": "y", "gvars": ["g1"], "xvars": ["x1"]}, "second": {"yvar": "y", "gvars": ["g1", "g2"], "xvars": ["x1", "x2"]}, "third": {"yvar": "y", "gvars": ["g2", "g3"], "xvars": ["x2", "x3"]}, } df.with_columns( pl.struct( ( pl.col(specs[specnm]["yvar"]) - pl.col(specs[specnm]["yvar"]).mean().over(specs[specnm]["gvars"]) ).abs(), *specs[specnm]["xvars"], ) .map_elements( partial( ols_fitted, yvar=specs[specnm]["yvar"], xvars=specs[specnm]["xvars"] ) ) .over("date", "id") .alias(f"fitted_{specnm}") for specnm in list(specs.keys()) )
运行代码时出现如下错误:
InvalidOperationError: window expression not allowed in aggregation
不清楚为何聚合上下文不支持over操作,希望了解该如何处理此问题,或有无系统化的替代实现方案。
问题原因
错误的核心是:你在pl.struct()内部使用了窗口函数(.over(specs[specnm]["gvars"])),而这个struct又被嵌套在另一个窗口上下文(.over("date", "id"))中,Polars不允许这种嵌套窗口表达式的写法。
修复方案:提前计算去均值后的y值
先把每个spec需要的去均值(组内中心化)后的y列计算出来,再构建struct进行回归拟合,避免嵌套窗口。
修正后的代码:
from functools import partial import polars as pl import numpy as np def ols_fitted(s: pl.Series, yvar: str, xvars: list[str]) -> pl.Series: df = s.struct.unnest() y = df[yvar].to_numpy() X = df[xvars].to_numpy() # 处理X列全0或共线性的情况,避免lstsq报错 try: coeffs = np.linalg.lstsq(X, y, rcond=None)[0] fitted = np.dot(X, coeffs) except np.linalg.LinAlgError: fitted = np.full_like(y, np.nan) return pl.Series(values=fitted, nan_to_null=True) df = pl.DataFrame( { "date": [1, 1, 1, 1, 2, 2, 2, 2, 2, 2], "id": [1, 1, 1, 2, 2, 2, 2, 3, 3, 3], "y": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10], "g1": [1, 1, 1, 1, 1, 2, 2, 2, 2, 2], "g2": [1, 1, 1, 2, 2, 2, 3, 3, 3, 3], "g3": [1, 1, 2, 2, 2, 3, 3, 4, 4, 4], "x1": [2, 5, 4, 7, 3, 2, 5, 6, 7, 2], "x2": [1, 5, 3, 4, 5, 6, 4, 3, 2, 1], "x3": [3, 6, 8, 6, 4, 7, 5, 4, 8, 1], } ) specs = { "first": {"yvar": "y", "gvars": ["g1"], "xvars": ["x1"]}, "second": {"yvar": "y", "gvars": ["g1", "g2"], "xvars": ["x1", "x2"]}, "third": {"yvar": "y", "gvars": ["g2", "g3"], "xvars": ["x2", "x3"]}, } # 第一步:提前计算每个spec对应的去均值y列 pre_processed = df.with_columns( ( pl.col(spec["yvar"]) - pl.col(spec["yvar"]).mean().over(spec["gvars"]) ).abs().alias(f"centered_y_{specnm}") for specnm, spec in specs.items() ) # 第二步:基于预处理后的列,按date+id分组拟合回归 result = pre_processed.with_columns( pl.struct( pl.col(f"centered_y_{specnm}"), *spec["xvars"], ) .map_elements( partial( ols_fitted, yvar=f"centered_y_{specnm}", xvars=spec["xvars"] ) ) .over("date", "id") .alias(f"fitted_{specnm}") for specnm, spec in specs.items() ) print(result)
优化建议:避免map_elements(提升性能)
map_elements是逐组遍历的Python级操作,数据量大时性能较差。可以改用Polars的向量化API结合group_by+apply来实现,性能更优:
def ols_fitted_group(df: pl.DataFrame, yvar: str, xvars: list[str]) -> pl.DataFrame: y = df[yvar].to_numpy() X = df[xvars].to_numpy() try: coeffs = np.linalg.lstsq(X, y, rcond=None)[0] fitted = np.dot(X, coeffs) except np.linalg.LinAlgError: fitted = np.full_like(y, np.nan) return df.with_columns(pl.Series(fitted, name=f"fitted")) # 预处理步骤同上 pre_processed = df.with_columns( ( pl.col(spec["yvar"]) - pl.col(spec["yvar"]).mean().over(spec["gvars"]) ).abs().alias(f"centered_y_{specnm}") for specnm, spec in specs.items() ) # 逐个处理每个spec的拟合 result = pre_processed for specnm, spec in specs.items(): result = result.group_by("date", "id", maintain_order=True).apply( lambda group: ols_fitted_group(group, f"centered_y_{specnm}", spec["xvars"]) ).rename({"fitted": f"fitted_{specnm}"}) print(result)
关键说明
- 避免嵌套窗口:窗口函数(
.over())不能嵌套在另一个窗口或聚合上下文里,必须提前计算出中间结果。 - 性能优化:
group_by.apply比map_elements.over()的性能更稳定,尤其是处理大数据集时,因为它直接操作分组后的DataFrame,减少了struct序列化/反序列化的开销。 - 异常处理:增加了线性代数异常的捕获,避免因X列共线性或全零导致的崩溃。
内容的提问来源于stack exchange,提问作者lebesgue
相关产品推荐
相关产品推荐

