如何从Pytorch Geometric的adj_t稀疏邻接矩阵获取edge_index张量
PyG中adj_t稀疏邻接矩阵转换为edge_index的方法
PyG内置的data.adj_t默认是转置后的稀疏邻接矩阵,类型为torch_sparse.SparseTensor,你可以直接调用其内置方法完成转换:
- 标准转换代码(适配PyG默认
adj_t存储格式):
# 取出转置后邻接矩阵的COO格式坐标 row, col, _ = data.adj_t.t().coo() # 拼接为形状[2, num_edges]的edge_index edge_index = torch.stack([row, col], dim=0)
注意:如果你的
adj_t是自定义生成、没有做过转置,存储的就是「源节点→目标节点」的邻接关系,可以去掉代码中的.t()调用,直接执行data.adj_t.coo()即可。
如果你的adj_t是scipy稀疏矩阵类型,可以用如下方法转换:
coo_adj = data.adj_t.tocoo() edge_index = torch.tensor([coo_adj.row, coo_adj.col], dtype=torch.long)
内容的提问来源于stack exchange,提问作者Qubix
相关产品推荐
相关产品推荐

