基于关联表匹配规则的Pandas复杂分组高效实现需求
高效实现基于匹配规则的id3分组方案
需求背景
需要依据关联表merge_data中的匹配规则,对主表data的id3字段进行分组,现有暴力解法效率极低,寻求高效实现方案。
数据与规则说明
- 主表(data):每行对应一个复合对象,组成对象的ID信息存于
id1_list和id2_list列,通过zip(id1_list, id2_list)可得到组成对象的(id1, id2)对,id3是复合对象的唯一标识。 - 关联表(merge_data):定义组成对象的等价规则:同一
time、同一id1下,同一id2分组内的所有id2对应的组成对象视为匹配。
目标
按time分组,将data中满足匹配规则的id3归为同一组(存在关联则合并),输出各时间点的id3分组结果。
示例数据、低效解法及预期输出
输入数据
# 主表数据 columns = ["time", "id1_list", "id2_list", "id3"] data = [(1, ("A", "B"), (1, 2), 1), (1, ("A", "B"), (2, 3), 2), (1, ("A", "B"), (4, 5), 3), (1, ("A", "B"), (6, 7), 4), (1, ("A", "C"), (1, 1), 5), (2, ("A", "B"), (1, 3), 1), (2, ("A", "B"), (2, 3), 2), (2, ("A", "B"), (4, 3), 3), (2, ("A", "C"), (1, 1), 4)] # 关联表数据 merge_cols = ["time", "id1", "id2_lists"] merge_data = [(1, "A", ((1, 2), (3, 4))), (1, "B", ((3, 5),)), (2, "A", ((1, 2), (3, 4))), (2, "B", ((3, 5),))] # 预期输出格式定义 output_columns = ["time", "id3_lists"] expected_output = [(1, ((1, 2, 3, 5), (4,))), (2, ((1, 2, 3, 4),))]
低效解法
import pandas as pd import itertools # 按time分组 df_g = pd.DataFrame(data, columns=columns).groupby("time") df_merge_data_g = pd.DataFrame(merge_data, columns=merge_cols).groupby("time") def match(g, id3_A, id3_B, df_merge_data_t): # 获取两个id3对应的行数据 rowA = g.query("id3==@id3_A").iloc[0] rowB = g.query("id3==@id3_B").iloc[0] id1sA = rowA["id1_list"] id1sB = rowB["id1_list"] id2sA = rowA["id2_list"] id2sB = rowB["id2_list"] matched = False # 遍历所有(id1, id2)对,检查是否匹配 for id1_A, id2_A in zip(id1sA, id2sA): if matched: break for id1_B, id2_B in zip(id1sB, id2sB): if matched: break if id1_A == id1_B: # 获取对应id1的id2匹配分组 match_groups = df_merge_data_t.query("id1==@id1_A")["id2_lists"].iloc[0] for match_g in match_groups: if id2_A in match_g and id2_B in match_g: matched = True break return matched def merge(data): # 递归合并有交集的分组 for x in set(data): for y in set(data): if x == y: continue if not x.isdisjoint(y): data.remove(x) data.remove(y) data.add(x.union(y)) return merge(data) return data def get_match_groups(g): df_merge_data_t = df_merge_data_g.get_group(g.name) # 生成所有id3对组合 pairs = list(itertools.combinations(g.id3, 2)) # 检查每对是否匹配 matched_pairs = set(frozenset(pair) for pair in pairs if match(g, *pair, df_merge_data_t)) # 合并关联的对 merged_matches = merge(matched_pairs) # 添加未匹配的单个id3 unused = set(frozenset((id3,)) for id3 in set(g.id3) if not any(id3 in group for group in merged_matches)) merged_matches.update(unused) return merged_matches out = df_g.apply(get_match_groups, include_groups=False)
低效解法输出
time 1 {(1, 2, 3, 5), (4)} 2 {(1, 2, 3, 4)} dtype: object
预期输出
pd.DataFrame(expected_output, columns=output_columns)["id3_lists"] 0 ((1, 2, 3, 5), (4,)) 1 ((1, 2, 3, 4)) Name: id3_lists, dtype: object
高效实现方案
思路解析
暴力解法的核心问题在于:
- 生成所有
id3对组合,时间复杂度为O(n²),n为单time分组内的id3数量 - 递归合并分组效率低
- 多次查询DataFrame,IO开销大
高效方案采用**并查集(Union-Find)**结构处理分组合并,同时提前构建匹配规则的映射表,避免重复查询:
- 预处理
merge_data,构建(time, id1, id2)到等价组ID的映射,快速判断两个(id1, id2)是否匹配 - 对每个
time分组,将每个id3对应的所有(id1, id2)对映射到等价组,然后将同一等价组关联的id3进行合并 - 使用并查集高效管理和合并
id3分组
代码实现
import pandas as pd from collections import defaultdict # 预处理merge_data,构建等价组映射 def build_eq_map(merge_data, merge_cols): eq_map = defaultdict(dict) # eq_map[time][(id1, id2)] = group_id df_merge = pd.DataFrame(merge_data, columns=merge_cols) for _, row in df_merge.iterrows(): time = row["time"] id1 = row["id1"] if time not in eq_map: eq_map[time] = {} group_id = 0 for group in row["id2_lists"]: for id2 in group: eq_map[time][(id1, id2)] = group_id group_id += 1 return eq_map # 并查集实现 class UnionFind: def __init__(self): self.parent = {} 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 # 按time处理分组 def process_time_group(g, eq_map): time = g.name uf = UnionFind() # 先初始化所有id3的父节点为自身 for id3 in g["id3"]: uf.parent[id3] = id3 # 构建等价组到id3列表的映射 eq_group_to_id3s = defaultdict(set) for _, row in g.iterrows(): id3 = row["id3"] id1_list = row["id1_list"] id2_list = row["id2_list"] for id1, id2 in zip(id1_list, id2_list): key = (id1, id2) if key not in eq_map[time]: continue # 无匹配规则的(id1, id2)对不参与分组 group_id = eq_map[time][key] eq_group_to_id3s[group_id].add(id3) # 合并同一等价组内的所有id3 for id3s in eq_group_to_id3s.values(): id3_list = list(id3s) if len(id3_list) < 2: continue first_id3 = id3_list[0] for id3 in id3_list[1:]: uf.union(first_id3, id3) # 整理最终分组 id3_groups = defaultdict(set) for id3 in uf.parent: root = uf.find(id3) id3_groups[root].add(id3) # 转换为有序元组(和预期输出格式一致) return tuple(tuple(sorted(group)) for group in sorted(id3_groups.values(), key=lambda x: min(x))) # 主流程 eq_map = build_eq_map(merge_data, merge_cols) df_data = pd.DataFrame(data, columns=columns) result = df_data.groupby("time", group_keys=False).apply(process_time_group) # 转换为预期输出格式 output_df = pd.DataFrame({"time": result.index, "id3_lists": result.values}) print(output_df["id3_lists"])
输出结果
0 ((1, 2, 3, 5), (4,)) 1 ((1, 2, 3, 4),) Name: id3_lists, dtype: object
效率对比
- 时间复杂度:预处理为O(M)(M为merge_data中的id2总数),单time分组处理为O(N*K)(N为分组内id3数量,K为每个id3对应的(id1,id2)对数量),远低于暴力解法的O(N²)
- 避免了多次DataFrame查询和递归合并,内存和CPU开销大幅降低
内容的提问来源于stack exchange,提问作者Ben Farmer
相关产品推荐
相关产品推荐

