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也遵循相同逻辑
- 每个symbol的时间序列按指定范围拆分,对应
内容的提问来源于stack exchange,提问作者JJ Fantini
相关产品推荐
相关产品推荐

