torch_geometric中如何从邻接矩阵提取edge_attr参数
从numpy邻接矩阵提取PyG框架
edge_attr参数的方法 edge_attr是PyG中存储边特征的参数,和edge_index里的边一一对应,形状一般为[num_edges, num_edge_features],从numpy稠密邻接矩阵提取不需要手动双层循环遍历,按以下逻辑实现即可:
- 核心思路是先定位邻接矩阵里所有非零值的位置(对应边的两个端点),再提取这些位置存储的数值作为边的原始属性,两种常用实现方式如下。
首先导入基础依赖,以你给出的邻接矩阵做示例:
import numpy as np import torch from torch_geometric.data import Data adj = np.array([[0, 1, 0],[1, 0, 0],[0, 0, 0]])
实现方式1:numpy原生实现(无额外依赖)
直接用np.where定位所有非零元素的坐标,再提取对应值:
# 拿到所有非零边的行、列索引(对应边的起点、终点) row, col = np.where(adj != 0) # 组装成PyG要求的edge_index,形状为[2, 边数],类型为torch.long edge_index = torch.tensor(np.stack([row, col]), dtype=torch.long) # 按坐标提取邻接矩阵里的非零值作为边属性,单值边权补最后一维符合PyG规范 edge_attr = torch.tensor(adj[row, col], dtype=torch.float).unsqueeze(dim=-1)
实现方式2:scipy稀疏格式转换(大规模邻接矩阵效率更高)
邻接矩阵尺寸很大时,转COO稀疏格式的处理速度更快:
from scipy.sparse import coo_matrix # 稠密邻接矩阵转COO稀疏格式 adj_sparse = coo_matrix(adj) # 提取边索引 edge_index = torch.tensor( np.stack([adj_sparse.row, adj_sparse.col]), dtype=torch.long ) # 提取边属性 edge_attr = torch.tensor( adj_sparse.data, dtype=torch.float ).unsqueeze(dim=-1)
你可以直接组装成PyG的Data对象验证结果:
data = Data(edge_index=edge_index, edge_attr=edge_attr) print(data) # 针对你给出的示例邻接矩阵,输出为 Data(edge_index=[2, 2], edge_attr=[2, 1]) # 对应无向图中节点0和节点1之间的双向边,每条边的特征值为1
常见场景调整
- 不需要维度补全:如果后续逻辑不需要
edge_attr带最后一维特征维度,可以去掉.unsqueeze(dim=-1),直接使用一维张量即可。 - 过滤自环:如果邻接矩阵对角线存在非零自环但不需要保留,可以加掩码过滤:
mask = row != col # scipy版本换成mask = adj_sparse.row != adj_sparse.col即可 edge_index = torch.tensor(np.stack([row[mask], col[mask]]), dtype=torch.long) edge_attr = torch.tensor(adj[row[mask], col[mask]], dtype=torch.float).unsqueeze(-1) - 自定义边特征:如果不需要把邻接矩阵的非零值作为特征,可以直接生成和边数等长的特征张量,比如全1张量
edge_attr = torch.ones(edge_index.size(1), 1)。 - 多维边特征:如果邻接矩阵每个位置存储的是多维边特征向量,拿到非零坐标后,直接索引对应位置的特征向量组装即可。
内容的提问来源于stack exchange,提问作者Behemdolg Fire
相关产品推荐
相关产品推荐

