基于路径图过滤DataFrame:筛选含至少3条边的连通子图
问题描述
给定如下DataFrame:
df = {'col1': ['a', 'b', 'c', 'd', 'e', 'w', 't', 'y', 'r', 's', 'n', 'm', 'p'], 'col2': ['b', 'c','d','e','f', 'x', 'z', 'w', 'w', 'n', 'm', 'p', 'q'], 'col3': [1, 2, 3, 4, 5, 6, 7, 8 ,9, 10, 11, 12, 13]}
对应的表格形式:
| 序号 | col1 | col2 | col3 |
|---|---|---|---|
| 0 | a | b | 1 |
| 1 | b | c | 2 |
| 2 | c | d | 3 |
| 3 | d | e | 4 |
| 4 | e | f | 5 |
| 5 | w | x | 6 |
| 6 | t | z | 7 |
| 7 | y | w | 8 |
| 8 | r | w | 9 |
| 9 | s | n | 10 |
| 10 | n | m | 11 |
| 11 | m | p | 12 |
| 12 | p | q | 13 |
需求:筛选出由col1和col2节点构成的线性路径子图(即每个中间节点仅有一个前驱和一个后继的连续链)中,包含至少3条边的子图对应的DataFrame行。预期输出如下:
| 序号 | col1 | col2 | col3 |
|---|---|---|---|
| 0 | a | b | 1 |
| 1 | b | c | 2 |
| 2 | c | d | 3 |
| 3 | d | e | 4 |
| 4 | e | f | 5 |
| 9 | s | n | 10 |
| 10 | n | m | 11 |
| 11 | m | p | 12 |
| 12 | p | q | 13 |
解决方案
通过节点的前驱/后继映射识别线性链,再筛选符合边数要求的链:
import pandas as pd # 构建DataFrame df = pd.DataFrame({ 'col1': ['a', 'b', 'c', 'd', 'e', 'w', 't', 'y', 'r', 's', 'n', 'm', 'p'], 'col2': ['b', 'c','d','e','f', 'x', 'z', 'w', 'w', 'n', 'm', 'p', 'q'], 'col3': [1, 2, 3, 4, 5, 6, 7, 8 ,9, 10, 11, 12, 13] }) # 创建节点的后继与前驱映射 successors = df.set_index('col1')['col2'].to_dict() predecessors = df.set_index('col2')['col1'].to_dict() visited_nodes = set() chain_id = 0 edge_chain_map = {} # 遍历所有节点,识别每条线性链 for node in df['col1']: if node not in visited_nodes: # 回溯找到链的起点 current_node = node while current_node in predecessors: prev_node = predecessors[current_node] if prev_node in visited_nodes: break current_node = prev_node if current_node in visited_nodes: continue # 顺向遍历整条链,记录所有边 current_chain_edges = [] while current_node in successors: next_node = successors[current_node] current_chain_edges.append((current_node, next_node)) visited_nodes.add(current_node) current_node = next_node visited_nodes.add(current_node) # 给当前链的所有边分配ID for edge in current_chain_edges: edge_chain_map[edge] = chain_id chain_id += 1 # 为每条边标记所属链ID df['chain_id'] = df.apply(lambda row: edge_chain_map.get((row['col1'], row['col2']), -1), axis=1) # 筛选出边数≥3的链 valid_chain_ids = df['chain_id'].value_counts()[lambda x: x >=3].index # 提取结果并移除辅助列 result = df[df['chain_id'].isin(valid_chain_ids)].drop('chain_id', axis=1) print(result)
运行后输出与预期一致。
内容的提问来源于stack exchange,提问作者Nico Rafael
相关产品推荐
相关产品推荐

