如何在Polars的两列中对所有存在关联的记录进行分组?
如何在Polars的两列中对所有存在关联的记录进行分组?
嘿,我来帮你搞定这个问题!你要的其实是把两列里有任何关联的记录都归到同一组——就像找社交网络里的好友链,只要能通过key1或key2连起来的,哪怕是间接关联,都算一伙的。比如你给的例子里,1号记录的key1是a,和2号的a关联;2号的key2是y,又和3号的y关联,所以这三条就属于同一组,4号没和任何其他记录沾边,就单独一组。
先明确你的输入数据,用Polars构造的话是这样:
import polars as pl df = pl.DataFrame( { "id": [1, 2, 3, 4], "key1": ["a", "a", "b", "c"], "key2": ["x", "y", "y", "z"], } )
对应的ASCII表格就是你给出的输入:
┌─────┬──────┬──────┐ │ id ┆ key1 ┆ key2 │ │ --- ┆ --- ┆ --- │ │ i64 ┆ str ┆ str │ ╞═════╪══════╪══════╡ │ 1 ┆ a ┆ x │ │ 2 ┆ a ┆ y │ │ 3 ┆ b ┆ y │ │ 4 ┆ c ┆ z │ └─────┴──────┴──────┘
解决思路
这个问题本质是找图的连通分量:把每个key(不管是key1还是key2里的)当成一个节点,每一行的key1和key2之间连一条线,最后所有能通过线连起来的节点就属于同一个组,对应的记录自然也归为一组。
我给你两种实现方式,一种是用第三方库快速搞定,另一种是纯Polars+Python实现,不用额外依赖。
方式一:用networkx快速实现(推荐)
networkx这个库处理图相关的问题超方便,几行代码就能搞定连通分量的计算:
import networkx as nx # 第一步:提取所有key1和key2的配对关系,作为图的边 edges = df.select(pl.col("key1"), pl.col("key2")).rows() # 第二步:构建图并找出每个key所属的连通分量 G = nx.Graph(edges) component_map = {} # 给每个连通分量选一个代表key,用来当整个组的标识 for component in nx.connected_components(G): group_label = next(iter(component)) # 就用分量里第一个key当组名 for key in component: component_map[key] = group_label # 第三步:把组标识映射回原DataFrame,给每条记录分配组 df_with_group = df.with_columns( group=pl.col("key1").map_dict(component_map) ) # 第四步:分组统计每组的记录数,得到你要的结果 result = df_with_group.group_by("group").agg( pl.count("id").alias("len") ).rename({"group": "key1 (or key2)"}) print(result)
运行后得到的结果就是你想要的:
┌────────────────┬─────┐ │ key1 (or key2) ┆ len │ │ --- ┆ --- │ │ str ┆ i64 │ ╞════════════════╪═════╡ │ a ┆ 3 │ │ c ┆ 1 │ └────────────────┴─────┘
方式二:纯Polars+Python实现(无额外依赖)
如果不想装第三方库,也可以用**并查集(Union-Find)**算法来实现,思路是不断合并关联的key,最终找出所有连通的组:
# 第一步:把所有key1和key2的配对关系整理成双向的长表 key_pairs = df.select(pl.col("key1").alias("key"), pl.col("key2").alias("pair")) key_pairs = key_pairs.vstack(df.select(pl.col("key2").alias("key"), pl.col("key1").alias("pair"))) # 第二步:初始化并查集,每个key的父节点先设为自己 parent = {key: key for key in key_pairs.select(pl.col("key")).unique().to_series().to_list()} # 实现并查集的查找函数(带路径压缩,提高效率) def find(u): while parent[u] != u: parent[u] = parent[parent[u]] u = parent[u] return u # 实现并查集的合并函数 def union(u, v): u_root = find(u) v_root = find(v) if u_root != v_root: parent[v_root] = u_root # 第三步:遍历所有key配对,合并关联的组 for key, pair in key_pairs.rows(): union(key, pair) # 第四步:生成每个key对应的组标识映射 component_map = {key: find(key) for key in parent.keys()} # 第五步:映射回原数据并分组统计 df_with_group = df.with_columns( group=pl.col("key1").map_dict(component_map) ) result = df_with_group.group_by("group").agg( pl.count("id").alias("len") ).rename({"group": "key1 (or key2)"}) print(result)
运行这个代码得到的结果和方式一完全一样,而且不需要任何额外的库,纯靠Polars和Python基础语法就能搞定。
不管用哪种方式,核心逻辑都是先找出所有关联key的集群,再把这个集群映射回原数据,这样就能把所有有间接关联的记录都归到同一组啦!
备注:内容来源于stack exchange,提问作者crazydragon777
相关产品推荐
相关产品推荐

