如何基于自定义比较函数为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
相关产品推荐
相关产品推荐

