如何在Polars Python中实现多参数滚动计算(滚动回归/相关)
在Polars中实现滚动相关/回归计算
核心结论
完全可以在Polars的滚动窗口操作中实现多列计算(如滚动相关系数、滚动回归),且无需将计算逻辑放到.agg()之外,直接在滚动聚合阶段完成更高效。
滚动相关系数实现示例
以你需要的滚动Spearman相关系数为例,直接在rolling().agg()中通过pl.struct打包多列窗口数据,配合map_elements执行计算:
import polars as pl from datetime import datetime from scipy.stats import spearmanr # 示例数据 df = pl.DataFrame( { "time": pl.date_range( start=datetime(2021, 1, 1), end=datetime(2022, 1, 1), interval="1d", eager=True ), "x": pl.int_range(0, 366, eager=True).shuffle(seed=1), "y": pl.int_range(0, 366, eager=True).shuffle(seed=2) } ) # 14天滚动Spearman相关系数 df_rolling_corr = df.rolling( index_column='time', period='14d' ).agg( pl.struct(["x", "y"]).map_elements( lambda struct: spearmanr(struct["x"], struct["y"])[0], return_dtype=pl.Float64 ).alias("rolling_spearman_corr") )
滚动回归最优实现
同样的思路,直接在滚动聚合阶段完成回归系数计算,避免生成不必要的中间列:
import statsmodels.api as sm def get_rolling_ols_coeff(struct: pl.Struct) -> float: y = struct["y"] # 为回归添加截距项(若不需要可移除sm.add_constant) x = sm.add_constant(struct["x"]) model = sm.OLS(y, x).fit() return model.params[1] # 返回自变量x的系数,params[0]为截距 df_rolling_reg = df.rolling( index_column='time', period='14d' ).agg( pl.struct(["x", "y"]).map_elements( get_rolling_ols_coeff, return_dtype=pl.Float64 ).alias("rolling_ols_coeff") )
你之前的错误原因
你尝试用pl.map_groups在rolling().agg()中操作是错误的:
map_groups是针对group_by后的分组数据设计的,而滚动窗口的agg中,每个表达式处理的是单个窗口内的列数据,并非分组后的整个数据集。- 正确的方式是用
pl.struct将窗口内的多列打包为结构体,再通过map_elements对每个窗口的结构体执行自定义计算。
临时方案的优化点
你的临时方案先聚合生成x_list、y_list中间列,再通过with_columns处理,会额外占用内存存储这些中间列。而直接在agg中用pl.struct+map_elements的方式,无需保留中间列,内存效率更高,尤其适合大规模数据集。
内容的提问来源于stack exchange,提问作者Scout
相关产品推荐
相关产品推荐

