基于torch_geometric实现K阶邻居的edge_index动态计算方法
计算PyTorch Geometric图的K阶邻居并生成新图
给定示例图的连通性:
edge_index = torch.tensor([[0, 1, 2, 3, 2, 0], [1, 0, 3, 2, 0, 2]], dtype=torch.long)
要实现一个函数,计算图的恰好K阶邻居(即通过恰好K步可达的节点对,排除小于K步的直接/间接邻居),并返回更新连通性后的新图edge_index。以下是实现方案:
实现思路
利用邻接矩阵的幂运算特性:邻接矩阵的K次幂中,非零元素对应K步可达的节点对。通过对比K次幂与K-1次幂的结果,筛选出仅K步可达的节点对,再转换回PyTorch Geometric的edge_index格式。
代码实现
import torch from torch_geometric.utils import to_adjacency_matrix, coalesce def get_k_hop_exact_neighbor_graph(edge_index, k, num_nodes=None): # 自动推断节点数量(若未指定) if num_nodes is None: num_nodes = edge_index.max().item() + 1 # 将edge_index转换为邻接矩阵,添加自环用于幂运算(代表0步可达自身) adj = to_adjacency_matrix(edge_index, num_nodes=num_nodes) adj_with_self = adj + torch.eye(num_nodes, device=adj.device) # 计算K步内可达的邻接矩阵 adj_k = torch.matrix_power(adj_with_self, k) if k > 1: # 计算K-1步内可达的邻接矩阵,用于排除小于K步的节点对 adj_k_minus_1 = torch.matrix_power(adj_with_self, k-1) # 保留仅K步可达的节点对,同时排除自环 exact_k_mask = (adj_k > 0) & (adj_k_minus_1 == 0) else: # K=1时,直接取原边(排除自环) exact_k_mask = adj > 0 # 将掩码转换为edge_index格式 edge_index_k = torch.nonzero(exact_k_mask, as_tuple=False).t().contiguous() # 去除重复边(确保无冗余的同方向边) edge_index_k, _ = coalesce(edge_index_k, None, num_nodes, num_nodes) return edge_index_k
测试示例
# 原示例图的edge_index edge_index = torch.tensor([[0, 1, 2, 3, 2, 0], [1, 0, 3, 2, 0, 2]], dtype=torch.long) # 计算2阶邻居图 edge_index_2hop = get_k_hop_exact_neighbor_graph(edge_index, k=2) print(edge_index_2hop)
输出结果(顺序可能因coalesce排序略有不同,但边集合与示例一致):
tensor([[0, 1, 1, 2, 3], [3, 2, 3, 1, 0]])
代码说明
- 邻接矩阵转换:
to_adjacency_matrix将edge_index转为稠密邻接矩阵,添加自环是为了让矩阵幂运算能正确计算步数(自环代表节点自身0步可达)。 - 幂运算筛选:通过对比K次幂和K-1次幂的邻接矩阵,只保留“K步可达但K-1步不可达”的节点对,确保是恰好K阶邻居。
- 格式转换与去重:将掩码矩阵转回
edge_index,并用coalesce去除重复边,保证输出格式符合PyTorch Geometric的要求。
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

