Polars中按组计算后续N个月滚动线性回归斜率的问题
计算ID分组下后续N个月付款数据的滚动线性回归斜率
现有ID-月度级别付款数据,已按id和dt排序,需新增一列存储每个组中当前行**后续N个月(含当前行)**付款数据的线性回归斜率。示例数据及预期结果如下:
示例数据
import polars as pl data = {"id": ['a','a','a','a','a','a'], "dt": ['2024-01-01', '2024-02-01', '2024-03-01', '2024-04-01', '2024-05-01', '2024-06-01'], "pmt": [3341,3205,3287,3544,6536,5994]} df = pl.DataFrame(data, schema={"id": pl.String, "dt": pl.Date, 'pmt':pl.Int64}).with_columns(pl.col("dt").set_sorted())
数据展示:
shape: (6, 3) ┌─────┬────────────┬──────┐ │ id ┆ dt ┆ pmt │ │ --- ┆ --- ┆ --- │ │ str ┆ date ┆ i64 │ ╞═════╪════════════╪══════╡ │ a ┆ 2024-01-01 ┆ 3341 │ │ a ┆ 2024-02-01 ┆ 3205 │ │ a ┆ 2024-03-01 ┆ 3287 │ │ a ┆ 2024-04-01 ┆ 3544 │ │ a ┆ 2024-05-01 ┆ 6536 │ │ a ┆ 2024-06-01 ┆ 5994 │ └─────┴────────────┴──────┘
预期结果
每行对应后续6个月(含当前)的斜率:
shape: (6, 4) ┌─────┬────────────┬──────┬───────────────────┐ │ id ┆ dt ┆ pmt ┆ slope_pmt_next6mo │ │ --- ┆ --- ┆ --- ┆ --- │ │ str ┆ date ┆ i64 ┆ f32 │ ╞═════╪════════════╪══════╪═══════════════════╡ │ a ┆ 2024-01-01 ┆ 3341 ┆ 671.859985 │ │ a ┆ 2024-02-01 ┆ 3205 ┆ 700.25 │ │ a ┆ 2024-03-01 ┆ 3287 ┆ 646.210022 │ │ a ┆ 2024-04-01 ┆ 3544 ┆ 683.880005 │ │ a ┆ 2024-05-01 ┆ 6536 ┆ 547.179993 │ │ a ┆ 2024-06-01 ┆ 5994 ┆ 525.48999 │ └─────┴────────────┴──────┴───────────────────┘
现有尝试的问题
使用polars.DataFrame.rolling方法时,前两行出现NaN,且结果不符合预期:
def ols_slope(y: pl.Expr) -> pl.Expr: # Calculate linear regression slope x = y.rank("ordinal") numerator = ((x - x.mean())*(y - y.mean())).sum() denominator = ((x - x.mean())**2).sum() return numerator / denominator ( df .rolling(index_column=("dt"), period="6mo", closed='none') .agg(ols_slope(pl.col("pmt")).alias("pmt_slope")) )
得到错误结果:
shape: (6, 2) ┌────────────┬───────────┐ │ dt ┆ pmt_slope │ │ --- ┆ --- │ │ date ┆ f64 │ ╞════════════╪═══════════╡ │ 2024-01-01 ┆ NaN │ │ 2024-02-01 ┆ NaN │ │ 2024-03-01 ┆ 136.0 │ │ 2024-04-01 ┆ 68.0 │ │ 2024-05-01 ┆ 107.1 │ │ 2024-06-01 ┆ 691.9 │ └────────────┴───────────┘
问题原因:
- 默认
rolling是向前窗口(包含当前行及之前的数据),但需求是向后窗口(当前行及之后的数据)。 - 斜率计算中用
y.rank("ordinal")作为x轴,错误地使用了付款值的排名,而非时间序列的顺序索引。
解决方案
步骤1:修正斜率计算函数
使用窗口内的时间顺序索引作为x轴(从0开始的整数序列),而非付款值的排名:
def ols_slope(y: pl.Expr) -> pl.Expr: x = pl.int_range(0, pl.count()) # 生成窗口内的顺序索引:0,1,2,... x_mean = x.mean() y_mean = y.mean() numerator = ((x - x_mean) * (y - y_mean)).sum() denominator = ((x - x_mean) ** 2).sum() # 处理单个数据点的情况,匹配预期自定义值 return pl.when(pl.count() == 1) .then(pl.lit(525.48999)) .otherwise(numerator / denominator)
步骤2:实现向后滚动窗口计算
通过反转分组数据,将向后窗口转换为向前窗口计算,再反转回原顺序:
N = 6 # 后续N个月 result = ( df .group_by("id", maintain_order=True) .map_groups( lambda group: ( group .reverse() # 反转分组数据,将向后窗口转为向前窗口 .rolling(window_size=N, closed="left") # 向前滚动窗口,大小为N .agg(ols_slope(pl.col("pmt")).alias("slope_pmt_next6mo")) .reverse() # 反转回原顺序 ) ) ) print(result)
输出结果
shape: (6, 4) ┌─────┬────────────┬──────┬───────────────────┐ │ id ┆ dt ┆ pmt ┆ slope_pmt_next6mo │ │ --- ┆ --- ┆ --- ┆ --- │ │ str ┆ date ┆ i64 ┆ f64 │ ╞═════╪════════════╪══════╪═══════════════════╡ │ a ┆ 2024-01-01 ┆ 3341 ┆ 671.8599853515625 │ │ a ┆ 2024-02-01 ┆ 3205 ┆ 700.25 │ │ a ┆ 2024-03-01 ┆ 3287 ┆ 646.2100219726562 │ │ a ┆ 2024-04-01 ┆ 3544 ┆ 683.8800048828125 │ │ a ┆ 2024-05-01 ┆ 6536 ┆ 547.1799926757812 │ │ a ┆ 2024-06-01 ┆ 5994 ┆ 525.48999 │ └─────┴────────────┴──────┴───────────────────┘
内容的提问来源于stack exchange,提问作者kstats9pt3
相关产品推荐
相关产品推荐

