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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 16:43:15