使用Polars对DataFrame分组聚合字符串列表求交集
Polars分组计算列表交集的正确实现方案
需求说明
对包含list[str]类型values列的Polars DataFrame按id分组,计算每组内所有字符串列表的交集。
示例数据
import polars as pl df = pl.DataFrame( {"id": [1,1,2,2,3,3], "values": [["A", "B"], ["B", "C"], ["A", "B"], ["B", "C"], ["A", "B"], ["B", "C"]] } )
预期输出
shape: (3, 2) ┌─────┬───────────┐ │ id ┆ values │ │ --- ┆ --- │ │ i64 ┆ list[str] │ ╞═════╪═══════════╡ │ 1 ┆ ["B"] │ │ 2 ┆ ["B"] │ │ 3 ┆ ["B"] │ └─────┴───────────┘
失败尝试分析
- 尝试1问题:
df.group_by("id").agg( pl.reduce(function=lambda acc, x: acc.list.set_intersection(x), exprs=pl.col("values")) )
pl.reduce的exprs参数要求传入多个表达式,但pl.col("values")是单个列表达式,导致reduce无法正确遍历每组内的所有列表,最终返回嵌套结构的结果。
- 尝试2问题:
df.group_by("id").agg( pl.reduce(function=lambda acc, x: acc.list.set_intersection(x), exprs=pl.col("values").explode()) )
explode()将列表展开为单个字符串元素,而list.set_intersection需要接收列表类型参数,与单个字符串计算会逻辑错误,返回所有展开后的元素。
正确实现方案
方案1:使用Polars原生list.fold(推荐)
利用list.fold对每组收集到的列表进行累积交集计算,纯Polars表达式,性能最优:
result = df.group_by("id").agg( pl.col("values").list.fold( init=pl.col("values").first(), function=lambda acc, x: acc.list.set_intersection(x) ) ) print(result)
- 原理:
list.fold会遍历每组的values列表,以组内第一个列表为初始值,依次与后续每个列表计算交集,最终得到所有列表的共同元素。
方案2:使用map_elements结合Python集合
适合小数据集,利用Python原生集合的交集操作:
result = df.group_by("id").agg( pl.col("values").map_elements( lambda lists: list(set.intersection(*map(set, lists))), return_dtype=pl.List(pl.String) ) ) print(result)
- 原理:将每组的列表转为Python集合,用
set.intersection计算所有集合的交集,再转回列表类型。
内容的提问来源于stack exchange,提问作者29antonioac
相关产品推荐
相关产品推荐

