如何在Polars中使用多列参数实现GroupBy Map UDF?
问题:Polars中多列参数调用Numba UDF异常的解决办法
我有一个Numba UDF:
@numba.jit(nopython=True) def generate_sample_numba(cumulative_dollar_volume: np.ndarray, dollar_tau: Union[int, np.ndarray]) -> np.ndarray: """ Generate the sample using numba for speed. """ covered_dollar_volume = 0 bar_index = 0 bar_index_array = np.zeros_like(cumulative_dollar_volume, dtype=np.uint32) if isinstance(dollar_tau, int): dollar_tau = np.array([dollar_tau] * len(cumulative_dollar_volume)) for i in range(len(cumulative_dollar_volume)): bar_index_array[i] = bar_index if cumulative_dollar_volume[i] >= covered_dollar_volume + dollar_tau[i]: bar_index += 1 covered_dollar_volume = cumulative_dollar_volume[i] return bar_index_array
该UDF接收两个输入:
cumulative_dollar_volumenumpy数组,对应group_by中的分组数据dollar_tau阈值,可以是整数或numpy数组
我重点关注numpy数组形式的dollar_tau,想要在Polars中实现与以下Pandas代码相同的效果:
data["bar_index"] = data.groupby(["ticker", "date"]).apply(lambda x: generate_sample_numba(x["cumulative_dollar_volume"].values, x["dollar_tau"].values)).explode().values.astype(int)
我尝试使用Polars的group_by().agg(pl.col().map_batches())实现:
cqt_sample = cqt_sample.with_columns( (pl.col("price") * pl.col("size")).alias("dollar_volume") ).with_columns( pl.col("dollar_volume").cum_sum().over(["ticker", "date"]).alias("cumulative_dollar_volume"), pl.lit(1_000_000).alias("dollar_tau") ) (cqt_sample .group_by(["ticker", "date"]) .agg(pl.col(["cumulative_dollar_volume", "dollar_tau"]) .map_batches(lambda x: generate_sample_numba(x["cumulative_dollar_volume"].to_numpy(), 1_000_000)) ) #.alias("bar_index") ) #.explode("bar_index")
但map_batches()返回了异常结果。不过当仅传入单列和整数dollar_tau时,代码运行正常:
(cqt_sample .group_by(["ticker", "date"]) .agg(pl.col("cumulative_dollar_volume") .map_batches(lambda x: generate_sample_numba(x.to_numpy(), 1_000_000)) ).alias("bar_index") ).explode("bar_index")
请问如何解决Polars中使用多列参数调用该Numba UDF时的异常问题?
解决方案
核心问题分析
当对多列使用map_batches时,传入lambda的x是一个Polars DataFrame(而非单列Series),但你的写法未正确处理多列提取逻辑,且测试代码硬编码了1_000_000作为dollar_tau,未实际使用分组后的dollar_tau列数据,这是异常的主要原因。
正确实现方式
需要在map_batches中从传入的DataFrame分别提取两列的numpy数组,传入Numba UDF后将结果包装回Polars Series,再展开匹配原数据行:
cqt_sample = cqt_sample.with_columns( (pl.col("price") * pl.col("size")).alias("dollar_volume") ).with_columns( pl.col("dollar_volume").cum_sum().over(["ticker", "date"]).alias("cumulative_dollar_volume"), pl.lit(1_000_000).alias("dollar_tau") ) # 多列参数调用逻辑 result = ( cqt_sample .group_by(["ticker", "date"]) .agg( pl.struct(["cumulative_dollar_volume", "dollar_tau"]) .map_batches( lambda df: pl.Series( generate_sample_numba( df["cumulative_dollar_volume"].to_numpy(), df["dollar_tau"].to_numpy() ), name="bar_index" ) ) ) .explode("bar_index") ) # 将结果合并回原表 cqt_sample = cqt_sample.join(result, on=["ticker", "date"])
关键改进点
- 使用
pl.struct打包需要的两列,确保map_batches能接收完整分组数据 - 在lambda中明确提取两列的numpy数组,传入Numba UDF
- 将UDF返回的数组包装为Polars Series并指定列名,方便后续
explode操作 - 通过
explode展开分组数组,再用join合并回原数据
可选优化:Polars 0.18+ 简洁写法
如果你的Polars版本≥0.18,可使用map方法替代map_batches,无需显式group_by和explode:
cqt_sample = cqt_sample.with_columns( (pl.col("price") * pl.col("size")).alias("dollar_volume") ).with_columns( cumulative_dollar_volume=pl.col("dollar_volume").cum_sum().over(["ticker", "date"]), dollar_tau=pl.lit(1_000_000) ).with_columns( bar_index=pl.struct(["cumulative_dollar_volume", "dollar_tau"]) .map(lambda s: generate_sample_numba(s["cumulative_dollar_volume"], s["dollar_tau"])) .over(["ticker", "date"]) )
内容的提问来源于stack exchange,提问作者Kevin Li
相关产品推荐
相关产品推荐

