如何通过检查两列集合的交集合并pandas DataFrame对应行
问题描述
我想要通过检查两列的集合交集来合并存储集合类型数据的DataFrame行。
我的DataFrame如下,每个单元格内的数据都是Python set() 类型:
import pandas as pd df = pd.DataFrame() df['COMPARE0'] = [set([1,2,3]),set([3,4,5]),set([6,7]),set([10,11]),set([12,13])] df['COMPARE1'] = [set(['a','b','c']),set(['d','e']),set(['c','f']),set(['g','h']),set(['h','i'])] df['GET0'] = [set(['aaa']),set(['bbb']),set(['ccc','ddd']),set(['efg','hii']),set(['efg','hiii'])] df['GET1'] = [set(['000']),set(['111']),set(['222']),set(['333','444']),set(['555'])]
输入样例

需求说明
基于COMPARE0和COMPARE1两列的集合交集结果合并所有行,只要任意一列的两个集合存在交集,就合并对应行的所有列数据。
期望输出

目前使用set()存储数据,也可以接受使用list类型的实现方案。
解决方法
这个需求本质是寻找行之间的连通分量:只要任意两行在COMPARE0或COMPARE1列存在集合交集,就属于同一分组,最终对同一分组的所有列做集合合并即可,以下是可直接运行的实现:
基础实现(适合小数据量)
# 并查集工具类 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): fx, fy = self.find(x), self.find(y) if fx != fy: self.parent[fy] = fx n = len(df) uf = UnionFind(n) # 遍历判断行之间是否需要合并 for i in range(n): for j in range(i + 1, n): has_intersection = bool(df.loc[i, 'COMPARE0'] & df.loc[j, 'COMPARE0']) or bool(df.loc[i, 'COMPARE1'] & df.loc[j, 'COMPARE1']) if has_intersection: uf.union(i, j) # 按分组合并所有列的集合 df['group'] = [uf.find(i) for i in range(n)] result = df.groupby('group').agg(lambda x: set.union(*x)).reset_index(drop=True)
优化实现(适合大数据量)
如果数据量较大,双层循环效率较低,可以通过元素映射的方式优化,时间复杂度接近线性:
from collections import defaultdict 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): fx, fy = self.find(x), self.find(y) if fx != fy: self.parent[fy] = fx n = len(df) uf = UnionFind(n) # 处理COMPARE0列,同元素对应的行直接合并 val_to_rows = defaultdict(list) for idx, s in enumerate(df['COMPARE0']): for val in s: val_to_rows[val].append(idx) for rows in val_to_rows.values(): for i in range(1, len(rows)): uf.union(rows[0], rows[i]) # 处理COMPARE1列,同元素对应的行直接合并 val_to_rows = defaultdict(list) for idx, s in enumerate(df['COMPARE1']): for val in s: val_to_rows[val].append(idx) for rows in val_to_rows.values(): for i in range(1, len(rows)): uf.union(rows[0], rows[i]) # 分组合并 df['group'] = [uf.find(i) for i in range(n)] result = df.groupby('group').agg(lambda x: set.union(*x)).reset_index(drop=True)
如果需要输出为list类型,只需要对最终结果的每一列调用list()方法转换即可。
内容的提问来源于stack exchange,提问作者Anonymous Anonymous
相关产品推荐
相关产品推荐

