能否在Polars中为search_sorted函数添加窗口功能?
Polars中实现分组窗口下的search_sorted功能
基础search_sorted使用示例
先看一个全局范围内的search_sorted用法:用df1的left_on值在df2的right_on序列中查找插入位置:
import polars as pl df1 = pl.DataFrame({ 'left_on': ['b', 'b', 'd', 'd'], }) df2 = pl.DataFrame({ 'right_on': ['a', 'b', 'c', 'd', 'e'], }) ( df1 .lazy() .with_context(df2.lazy()) .with_columns([ pl.col('right_on').search_sorted(pl.col('left_on'), side='left') ]) .collect() )
执行结果:
┌─────────┬──────────┐ │ left_on ┆ right_on │ │ --- ┆ --- │ │ str ┆ u32 │ ╞═════════╪══════════╡ │ b ┆ 1 │ │ b ┆ 1 │ │ d ┆ 3 │ │ d ┆ 3 │ └─────────┴──────────┘
需求:分组维度下的search_sorted
现在需要实现分组窗口内的search_sorted:df1和df2分别有分组键by_left和by_right,要求仅在同分组内,用df1的left_on值查找df2对应分组right_on序列的插入位置,期望结果如下:
import polars as pl df1 = pl.DataFrame({ 'left_on': ['b', 'b', 'd', 'd'] * 2, 'by_left': ['X'] * 4 + ['Y'] * 4, }) df2 = pl.DataFrame({ 'right_on': ['a', 'b', 'c', 'd', 'e'] * 2, 'by_right': ['X'] * 5 + ['Y'] * 5, }) # 期望输出: """ ┌─────────┬─────────┬──────────┐ │ left_on ┆ by_left ┆ right_on │ │ --- ┆ --- ┆ --- │ │ str ┆ str ┆ u32 │ ╞═════════╪═════════╪══════════╡ │ b ┆ X ┆ 1 │ │ b ┆ X ┆ 1 │ │ d ┆ X ┆ 3 │ │ d ┆ X ┆ 3 │ │ b ┆ Y ┆ 6 │ │ b ┆ Y ┆ 6 │ │ d ┆ Y ┆ 8 │ │ d ┆ Y ┆ 8 │ └─────────┴─────────┴──────────┘ """
为什么不用join_asof?
join_asof虽能实现类似关联,但有两个关键局限:
- 不支持
search_sorted的side参数(any/left/right)完整选项; - 无法对UTF-8字符串类型执行搜索。
尝试过的无效方法
直接在search_sorted后追加.over('by_left')无法得到有效结果:
pl.col('right_on').search_sorted(pl.col('left_on'), side='left').over('by_left')
解决方案
通过group_by+map_groups实现分组内的search_sorted:对每个分组单独关联对应分组的df2数据,再执行搜索:
import polars as pl df1 = pl.DataFrame({ 'left_on': ['b', 'b', 'd', 'd'] * 2, 'by_left': ['X'] * 4 + ['Y'] * 4, }) df2 = pl.DataFrame({ 'right_on': ['a', 'b', 'c', 'd', 'e'] * 2, 'by_right': ['X'] * 5 + ['Y'] * 5, }) result = ( df1.lazy() .group_by('by_left', maintain_order=True) .map_groups(lambda group: group.with_context( df2.lazy().filter(pl.col('by_right') == group['by_left'].first()) ).with_columns( search_idx=pl.col('right_on').search_sorted(pl.col('left_on'), side='left') ) ) .collect() ) print(result)
执行后即可得到期望的分组搜索结果。
内容的提问来源于stack exchange,提问作者T.H Rice
相关产品推荐
相关产品推荐

