在Python Polars中计算时如何获取前一个计算值?
用Polars实现自定义递推计算的方案
你的递推公式 result_col[i] = s[i] * c[i] + (1 - s[i]) * result_col[i-1] 与Polars中ewm_mean(adjust=False)的递归计算逻辑完全一致,因此可以直接利用Polars的内置函数高效实现,无需手动写循环。
直接用ewm_mean实现
Polars的ewm_mean函数在adjust=False时,采用的递归公式为:
y[t] = α * x[t] + (1 - α) * y[t-1]
其中α对应你的s[i],x[t]对应你的c[i],完全匹配你的需求。即使s[i]超出常规EWM的0<α<1范围(比如你示例中的s=2),Polars仍会按照递归式正常计算,不会限制α的取值范围。
代码示例:
import polars as pl # 构造示例数据 df = pl.DataFrame({ "s": [1, 2], "c": [3, 5] }) # 计算递推列 df = df.with_columns( result_col=pl.col("c").ewm_mean(alpha=pl.col("s"), adjust=False) ) print(df)
运行结果:
shape: (2, 3) ┌─────┬─────┬────────────┐ │ s ┆ c ┆ result_col │ │ --- ┆ --- ┆ --- │ │ i64 ┆ i64 ┆ f64 │ ╞═════╪═════╪════════════╡ │ 1 ┆ 3 ┆ 3.0 │ │ 2 ┆ 5 ┆ 7.0 │ └─────┴─────┴────────────┘
结果完全符合手动计算的预期:
- 第一行:
1*3 + (1-1)*0 = 3(默认取第一个c值作为初始结果,无前置元素时自动适配) - 第二行:
2*5 + (1-2)*3 = 10 - 3 =7
自定义初始值的处理
如果你的递推需要指定第一个元素的初始值(比如不是默认的s[0]*c[0]),可以结合when/then手动设置:
# 示例:将第一个元素的初始值设为0 df = df.with_columns( idx=pl.int_range(0, pl.count()), temp_ewm=pl.col("c").ewm_mean(alpha=pl.col("s"), adjust=False) ).with_columns( result_col=pl.when(pl.col("idx") == 0) .then(pl.col("s")*pl.col("c") + (1-pl.col("s"))*0) # 自定义初始值 .otherwise(pl.col("temp_ewm")) ).drop("idx", "temp_ewm")
若无法使用ewm_mean的优化方案
如果后续递推逻辑发生变化,无法通过内置函数覆盖,可采用以下优化方案:
1. map_batches结合NumPy向量化
将数据按批次处理,用NumPy的向量化操作替代原生Python循环,大幅提升速度:
import polars as pl import numpy as np def custom_recursion(batch: pl.DataFrame) -> pl.Series: s = batch["s"].to_numpy() c = batch["c"].to_numpy() result = np.zeros_like(c, dtype=np.float64) result[0] = s[0] * c[0] for i in range(1, len(result)): result[i] = s[i] * c[i] + (1 - s[i]) * result[i-1] return pl.Series(result, name="result_col") df = df.with_columns( pl.struct(["s", "c"]).map_batches(custom_recursion) )
2. Numba编译加速循环
如果递推逻辑必须用循环,可借助Numba将Python函数编译为机器码,进一步提升性能:
import polars as pl import numpy as np from numba import jit @jit(nopython=True) def numba_recursion(s: np.ndarray, c: np.ndarray) -> np.ndarray: result = np.zeros_like(s, dtype=np.float64) result[0] = s[0] * c[0] for i in range(1, len(result)): result[i] = s[i] * c[i] + (1 - s[i]) * result[i-1] return result df = df.with_columns( result_col=pl.col("s", "c").map_batches(lambda x: numba_recursion(x["s"].to_numpy(), x["c"].to_numpy())) )
3. 分组递推用fold操作
如果需要按分组执行递推计算,可使用Polars的fold操作结合窗口函数实现分组内的递推:
# 替换"your_group_column"为实际分组列名 df = df.with_columns( result_col=pl.fold( acc=pl.lit(0.0), function=lambda acc, x: x["s"] * x["c"] + (1 - x["s"]) * acc, exprs=pl.struct(["s", "c"]) ).over("your_group_column") )
内容的提问来源于stack exchange,提问作者Göksel Gürsel
相关产品推荐
相关产品推荐

