You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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_volume numpy数组,对应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"])

关键改进点

  1. 使用pl.struct打包需要的两列,确保map_batches能接收完整分组数据
  2. 在lambda中明确提取两列的numpy数组,传入Numba UDF
  3. 将UDF返回的数组包装为Polars Series并指定列名,方便后续explode操作
  4. 通过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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 08:50:37