Polars升级后arr.get()报错,求替代方案与优化方法
Polars 0.18.4版本适配与线性回归多列优化方案
一、修复arr.get()报错
Polars 0.18+版本移除了数组的get()方法,替换为以下两种写法:
写法1:使用arr.element()方法
df.select([ pl.all().exclude("elapsed_time_linreg"), pl.col("elapsed_time_linreg").arr.element(0).suffix("_slope"), pl.col("elapsed_time_linreg").arr.element(1).suffix("_intercept"), pl.col("elapsed_time_linreg").arr.element(2).suffix("_resid_std"), ])
写法2:直接下标索引
df.select([ pl.all().exclude("elapsed_time_linreg"), pl.col("elapsed_time_linreg")[0].suffix("_slope"), pl.col("elapsed_time_linreg")[1].suffix("_intercept"), pl.col("elapsed_time_linreg")[2].suffix("_resid_std"), ])
二、更高效的多列生成方式
原方案先生成数组列再拆分存在额外开销,推荐直接生成目标列:
方法1:返回元组+Struct展开(保留numba加速)
修改函数返回元组,通过Polars的Struct类型直接展开为多列,避免数组列转换:
from numba import jit import numpy as np import polars as pl @jit def linear_regression(session_np: np.ndarray) -> tuple[float, float, float]: w = len(session_np) x = np.arange(w) sx = w ** 2 / 2 sy = np.sum(session_np) sx2 = (w * (w + 1) * (2 * w + 1)) / 6 sxy = np.sum(x * session_np) slope = (w * sxy - sx * sy) / (w * sx2 - sx**2) intercept = (sy - slope * sx) / w resids = session_np - (x * slope + intercept) return slope, intercept, resids.std() def get_linreg_aggs(session) -> tuple[float, float, float]: return linear_regression(np.array(session)) # 直接生成多列 df = df.with_columns([ pl.col("elapsed_time") .apply(get_linreg_aggs, return_dtype=pl.Struct([ pl.Field("slope", pl.Float32), pl.Field("intercept", pl.Float32), pl.Field("resid_std", pl.Float32) ])) .struct.unnest() .suffix("_elapsed_time") ])
方法2:Polars原生表达式实现(无Python层开销)
完全用Polars原生函数替代apply,性能更优:
df = df.with_columns([ # 预计算线性回归参数 pl.col("elapsed_time").map_elements(len, return_dtype=pl.Int64).alias("w"), pl.col("elapsed_time").sum().alias("sy"), pl.col("elapsed_time").map_elements(lambda s: np.sum(np.arange(len(s)) * s), return_dtype=pl.Float32).alias("sxy"), ]).with_columns([ (pl.col("w").pow(2) / 2).alias("sx"), (pl.col("w") * (pl.col("w") + 1) * (2 * pl.col("w") + 1) / 6).alias("sx2"), ]).with_columns([ # 计算最终指标 ((pl.col("w") * pl.col("sxy") - pl.col("sx") * pl.col("sy")) / (pl.col("w") * pl.col("sx2") - pl.col("sx").pow(2))).alias("slope"), ((pl.col("sy") - pl.col("slope") * pl.col("sx")) / pl.col("w")).alias("intercept"), pl.col("elapsed_time").map_elements( lambda s, slope, intercept: (s - (np.arange(len(s)) * slope + intercept)).std(), args=[pl.col("slope"), pl.col("intercept")], return_dtype=pl.Float32 ).alias("resid_std"), ]).drop(["w", "sy", "sxy", "sx", "sx2"])
三、关于lru-cache在Polars多进程环境的有效性
- Polars默认多进程执行时,
lru-cache是进程内缓存,不同进程间不共享缓存,每个进程会单独维护一份缓存副本,导致内存占用增加,缓存命中率下降。 - 如果输入数据重复率极高,可改用进程间共享缓存方案(如基于
multiprocessing.Manager实现的共享缓存),但会增加代码复杂度;若重复率低,不建议使用,收益有限。
内容的提问来源于stack exchange,提问作者yk4r2
相关产品推荐
相关产品推荐

