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

如何基于自定义比较函数为Pandas DataFrame生成分组ID?

问题描述

我有一个用于比较DataFrame行的函数:

def comp(lhs: pandas.Series, rhs: pandas.Series) -> bool:
    if lhs.id == rhs.id:
        return True

    if abs(lhs.val1 - rhs.val1) < 1e-8:
        if abs(lhs.val2 - rhs.val2) < 1e-8:
            return True

    return False

现在我有一个包含id、val1和val2列的DataFrame,希望生成分组ID,使得任意两个经comp函数判断为True的行拥有相同分组编号。尝试用groupby但没找到合适方法。

最小可复现示例(MRE):

import pandas as pd

example_input = pd.DataFrame({
    'id' : [0, 1, 2, 2, 3],
    'val1' : [1.1, 1.2, 1.3, 1.4, 1.1],
    'val2' : [2.1, 2.2, 2.3, 2.4, 2.1]
})

# 期望输出
example_output = example_input.copy()
example_output.index = [0, 1, 2, 2, 0]
example_output.index.name = 'groups'
解决方案

这个问题本质是连通分量匹配:满足comp条件的行属于同一连通组,groupby无法处理这种间接关联的分组,需要用并查集(Union-Find)算法实现。

实现步骤

1. 定义并查集工具类

用于高效管理和合并连通组:

class UnionFind:
    def __init__(self, size):
        self.parent = list(range(size))
    
    def find(self, x):
        # 路径压缩优化
        if self.parent[x] != x:
            self.parent[x] = self.find(self.parent[x])
        return self.parent[x]
    
    def union(self, x, y):
        # 合并两个连通分量
        x_root = self.find(x)
        y_root = self.find(y)
        if x_root != y_root:
            self.parent[y_root] = x_root

2. 分阶段合并连通组

按照comp函数的逻辑,分两步合并行的关联关系:

  • 第一步:合并所有id相同的行
  • 第二步:合并所有val1和val2近似相等的行

3. 生成最终分组ID

将并查集中的根节点映射为唯一分组编号,添加到原DataFrame。

完整代码

import pandas as pd

def comp(lhs: pd.Series, rhs: pd.Series) -> bool:
    if lhs.id == rhs.id:
        return True
    if abs(lhs.val1 - rhs.val1) < 1e-8 and abs(lhs.val2 - rhs.val2) < 1e-8:
        return True
    return False

class UnionFind:
    def __init__(self, size):
        self.parent = list(range(size))
    
    def find(self, x):
        if self.parent[x] != x:
            self.parent[x] = self.find(self.parent[x])
        return self.parent[x]
    
    def union(self, x, y):
        x_root = self.find(x)
        y_root = self.find(y)
        if x_root != y_root:
            self.parent[y_root] = x_root

# 处理输入数据
example_input = pd.DataFrame({
    'id' : [0, 1, 2, 2, 3],
    'val1' : [1.1, 1.2, 1.3, 1.4, 1.1],
    'val2' : [2.1, 2.2, 2.3, 2.4, 2.1]
})

# 初始化并查集
uf = UnionFind(len(example_input))

# 合并同一id的行
for _, group in example_input.groupby('id'):
    indices = group.index.tolist()
    # 只需将组内元素与第一个元素合并,无需两两组合,优化效率
    base_idx = indices[0]
    for idx in indices[1:]:
        uf.union(base_idx, idx)

# 合并val1/val2近似相等的行
# 用乘以1e8取整的方式处理浮点数精度问题,与comp函数阈值匹配
example_input['val1_round'] = (example_input['val1'] * 1e8).astype(int)
example_input['val2_round'] = (example_input['val2'] * 1e8).astype(int)

for _, group in example_input.groupby(['val1_round', 'val2_round']):
    indices = group.index.tolist()
    base_idx = indices[0]
    for idx in indices[1:]:
        uf.union(base_idx, idx)

# 映射根节点为分组ID
root_to_group = {}
current_group = 0
groups = []
for idx in range(len(example_input)):
    root = uf.find(idx)
    if root not in root_to_group:
        root_to_group[root] = current_group
        current_group += 1
    groups.append(root_to_group[root])

# 生成期望格式的输出
example_output = example_input.drop(['val1_round', 'val2_round'], axis=1)
example_output.index = groups
example_output.index.name = 'groups'

print(example_output)

输出结果

id  val1  val2
groups                
0        0   1.1   2.1
1        1   1.2   2.2
2        2   1.3   2.3
2        2   1.4   2.4
0        3   1.1   2.1

优化说明

  • 并查集的路径压缩优化保证了接近O(n)的时间复杂度,适合处理大体积DataFrame
  • 浮点数近似处理采用与comp函数一致的阈值逻辑,避免精度误差导致的分组错误
  • 合并组时仅将组内元素与第一个元素合并,减少不必要的操作次数

内容的提问来源于stack exchange,提问作者quant

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 23:21:35