如何在Polars LazyFrame中用map_batches多列应用带参自定义函数?
在Polars LazyFrame中实现带次要列参数的多列自定义函数应用
我需要把已掌握的Pandas功能迁移到Polars LazyFrame:将自定义函数应用到多列,且该函数需要以另一列作为次要参数。
Pandas实现示例
import statsmodels.stats.proportion as ssp import pandas as pd df = pd.DataFrame( { "a": [0, 1, 0, 1], "b": [1, 2, 3, 4], "nobs": [10, 20, 30, 40], } ) proportion_confint_func = lambda c: ssp.proportion_confint(c, df["nobs"], method="agresti_coull")[0] df[["a", "b"]].apply(proportion_confint_func)
执行结果:
a b 0 0.0 0.000000 1 0.0 0.015656 2 0.0 0.026639 3 0.0 0.033880
我的Polars尝试及问题
我尝试用map_batches实现,但得到了object类型列,代码如下:
from functools import partial from scipy import stats import polars as pl def proportion_confint_pl(count_a, nobs_a, alpha:float=0.05, method="agresti_coull", bound="lower"): crit = stats.norm.isf(alpha / 2.0) nobs_c = nobs_a + crit**2 q_c = (count_a + crit**2 / 2.0) / nobs_c std_c = ((q_c * (1.0 - q_c) / nobs_c)).sqrt() dist = crit * std_c if bound == "lower": ci = (q_c - dist) elif bound == "upper": ci = (q_c + dist) return ci.clip(0, 1) ldf = pl.LazyFrame( { "a": [0, 1, 0, 1], "b": [1, 2, 3, 4], "nobs": [10, 20, 30, 40], } ) proportion_confint_pl_func = partial(proportion_confint_pl, nobs_a=pl.col("nobs")) ldf.select(pl.col(["a", "b"]).map_batches(proportion_confint_pl_func)).collect()
执行输出:
shape: (1, 2) ┌─────────────────────────────────┬─────────────────────────────────┐ │ a ┆ b │ │ --- ┆ --- │ │ object ┆ object │ ╞═════════════════════════════════╪═════════════════════════════════╡ │ [([(Series[a]) / ([(col("nobs"… ┆ [([(Series[b]) / ([(col("nobs"… │ └─────────────────────────────────┴─────────────────────────────────┘
请问该如何正确实现需求?
正确实现方法
方法1:使用pl.struct结合apply(推荐,适配LazyFrame)
将目标列与次要参数列打包成结构体,逐行传递给自定义函数,逻辑清晰且适配LazyFrame的延迟执行:
from scipy import stats import polars as pl def proportion_confint_pl(row, alpha:float=0.05, bound="lower"): count_a = row[0] nobs_a = row[1] crit = stats.norm.isf(alpha / 2.0) nobs_c = nobs_a + crit**2 q_c = (count_a + crit**2 / 2.0) / nobs_c std_c = ((q_c * (1.0 - q_c) / nobs_c))**0.5 dist = crit * std_c ci = q_c - dist if bound == "lower" else q_c + dist return max(0.0, min(1.0, ci)) ldf = pl.LazyFrame( { "a": [0, 1, 0, 1], "b": [1, 2, 3, 4], "nobs": [10, 20, 30, 40], } ) result = ldf.select( pl.struct(["a", "nobs"]).apply(proportion_confint_pl).alias("a"), pl.struct(["b", "nobs"]).apply(proportion_confint_pl).alias("b") ).collect() print(result)
输出结果:
shape: (4, 2) ┌──────────┬──────────┐ │ a ┆ b │ │ --- ┆ --- │ │ f64 ┆ f64 │ ╞══════════╪══════════╡ │ 0.0 ┆ 0.0 │ │ 0.0 ┆ 0.015656 │ │ 0.0 ┆ 0.026639 │ │ 0.0 ┆ 0.03388 │ └──────────┴──────────┘
方法2:修正map_batches的用法
map_batches的函数接收的是整个批次的DataFrame,而非单个列表达式,调整函数逻辑处理批次数据:
from scipy import stats import polars as pl def proportion_confint_batch(batch_df, alpha:float=0.05, bound="lower"): crit = stats.norm.isf(alpha / 2.0) nobs_c = batch_df["nobs"] + crit**2 result = {} for col_name in ["a", "b"]: count_a = batch_df[col_name] q_c = (count_a + crit**2 / 2.0) / nobs_c std_c = ((q_c * (1.0 - q_c) / nobs_c))**0.5 dist = crit * std_c ci = q_c - dist if bound == "lower" else q_c + dist result[col_name] = ci.clip(0, 1) return pl.DataFrame(result) ldf = pl.LazyFrame( { "a": [0, 1, 0, 1], "b": [1, 2, 3, 4], "nobs": [10, 20, 30, 40], } ) result = ldf.map_batches(proportion_confint_batch).collect() print(result)
输出结果与方法1一致。
方法3:用Polars原生表达式重构(最优,性能最佳)
将自定义逻辑转换成Polars链式表达式,完全利用LazyFrame的查询优化,适合大数据场景:
import polars as pl from scipy import stats alpha = 0.05 crit = stats.norm.isf(alpha / 2.0) ldf = pl.LazyFrame( { "a": [0, 1, 0, 1], "b": [1, 2, 3, 4], "nobs": [10, 20, 30, 40], } ) result = ldf.select( ( ((pl.col("a") + crit**2 / 2) / (pl.col("nobs") + crit**2)) - crit * pl.sqrt( ((pl.col("a") + crit**2 / 2) / (pl.col("nobs") + crit**2)) * (1 - ((pl.col("a") + crit**2 / 2) / (pl.col("nobs") + crit**2))) / (pl.col("nobs") + crit**2) ) ).clip(0, 1).alias("a"), ( ((pl.col("b") + crit**2 / 2) / (pl.col("nobs") + crit**2)) - crit * pl.sqrt( ((pl.col("b") + crit**2 / 2) / (pl.col("nobs") + crit**2)) * (1 - ((pl.col("b") + crit**2 / 2) / (pl.col("nobs") + crit**2))) / (pl.col("nobs") + crit**2) ) ).clip(0, 1).alias("b") ).collect() print(result)
输出符合预期,且性能最优。
内容的提问来源于stack exchange,提问作者Kyle Gilde
相关产品推荐
相关产品推荐

