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

Polars中优化拆分数据分组的索引分配方案

已解决:

最优函数:比原函数快995倍

def add_range_index_stack2(data, range_str):
    range_str = _range_format(range_str)
    df_range_index = (
        data.group_by_dynamic(index_column="date", every=range_str, by="symbol")
        .agg()
        .with_columns(
            pl.int_range(0, pl.len()).over("symbol").alias("range_index")
        )
    )
    data = data.join_asof(df_range_index, on="date", by="symbol")

    return data
原始问题:

数据逻辑:

  • 我有一组需要拆分为多个块的时间序列数据,以股票报价和每日价格数据为例:
    • 若时间序列时长3个月,拆分范围为1个月,则生成3个数据块,每个月份对应递增的整数标签
    • 新增range_index列,从0开始递增,比如1月数据标签为0、2月为1、3月为2
  • 需为DataFrame中的每个symbol独立执行此操作,不同symbol的start_date可能不同,要根据各自的股票数据准确分配range_index值

已实现方案:

我用Polars编写了一个函数来添加该列,但处理多symbol的多年数据时,执行时间慢到约3秒,希望找到更快的实现方式。我知道Polars中行式操作比列式操作慢,求Polars资深用户帮忙优化。

def add_range_index(
    data: pl.LazyFrame | pl.DataFrame, range_str: str
) -> pl.LazyFrame | pl.DataFrame:
    """
    Context: Toolbox || Category: Helpers || Sub-Category: Mandelbrot Channel Helpers || **Command: add_n_range**.

    This function is used to add a column to the dataframe that contains the
    range grouping for the entire time series.


    This function is used in `log_mean()`
    """  # noqa: W505
    range_str = _range_format(range_str)

    if "date" in data.columns:
        group_by_args = {
            "every": range_str,
            "closed": "left",
            "include_boundaries": True,
        }
        if "symbol" in data.columns:
            group_by_args["by"] = "symbol"
            symbols = (data.select("symbol").unique().count())["symbol"][0]

        grouped_data = (
            data.lazy()
            .set_sorted("date")
            .group_by_dynamic("date", **group_by_args)
            .agg(
                pl.col("adj_close").count().alias("n_obs")
            )  # using 'adj_close' as the column to sum
        )
    range_row = grouped_data.with_columns(
        pl.arange(0, pl.count()).over("symbol").alias("range_index")
    )
    ## WIP:
    # Extract the number of ranges the time series has

    # Initialize a new column to store the range index
    data = data.with_columns(pl.lit(None).alias("range_index"))

    # Loop through each range and add the range index to the original dataframe
    for row in range_row.collect().to_dicts():
        symbol = row["symbol"]
        start_date = row["_lower_boundary"]
        end_date = row["_upper_boundary"]
        range_index = row["range_index"]

        # Apply the conditional logic to each group defined by the 'symbol' column
        data = data.with_columns(
            pl.when(
                (pl.col("date") >= start_date)
                & (pl.col("date") < end_date)
                & (pl.col("symbol") == symbol)
            )
            .then(range_index)
            .otherwise(pl.col("range_index"))
            .over("symbol")  # Apply the logic over each 'symbol' group
            .alias("range_index")
        )

    return data


def _range_format(range_str: str) -> str:
    """
    Context: Toolbox || Category: Technical || Sub-Category: Mandelbrot Channel Helpers || **Command: _range_format**.

    This function formats a range string into a standard format.
    The return value is to be passed to `_range_days()`.

    Parameters
    ----------
    range_str : str
        The range string to format. It should contain a number followed by a
        range part. The range part can be 'day', 'week', 'month', 'quarter', or
        'year'. The range part can be in singular or plural form and can be
        abbreviated. For example, '2 weeks', '2week', '2wks', '2wk', '2w' are
        all valid.

    Returns
    -------
    str
        The formatted range string. The number is followed by an abbreviation of
        the range part ('d' for day, 'w' for week, 'mo' for month, 'q' for
        quarter, 'y' for year). For example, '2 weeks' is formatted as '2w'.

    Raises
    ------
    RangeFormatError
        If an invalid range part is provided.

    Notes
    -----
    This function is used in `log_mean()`
    """  # noqa: W505
    # Separate the number and range part
    num = "".join(filter(str.isdigit, range_str))

    # Find the first character after the number in the range string
    range_part = next((char for char in range_str if char.isalpha()), None)

    # Check if the range part is a valid abbreviation
    if range_part not in {"d", "w", "m", "y", "q"}:
        msg = f"`{range_str}` could not be formatted; needs to include d, w, m, y, q"
        raise HumblDataError(msg)

    # If the range part is "m", replace it with "mo" to represent "month"
    if range_part == "m":
        range_part = "mo"

    # Return the formatted range string
    return num + range_part

预期数据格式:

  • 新增range_index列后的数据需满足:
    • 每个symbol的时间序列按指定范围拆分,对应range_index从0开始递增
    • 不同symbol的拆分独立进行,各自的range_index从0开始计数
    • 示例:某symbol的1-3月数据,1月行的range_index为0,2月为1,3月为2;PCT等其他symbol也遵循相同逻辑

内容的提问来源于stack exchange,提问作者JJ Fantini

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 18:08:09