如何提取并可视化PyTorch Geometric模型的最后一层节点嵌入
问题
我正在使用torch_geometric开展首个图卷积神经网络(GCN)项目,已基于CiteSeer数据集成功训练了一个两层GCN模型,但不知如何提取模型的最后一层节点嵌入(emb)以及节点类型(node_type),以使用给定的visualize函数完成可视化。
数据集加载代码:
from torch_geometric.datasets import Planetoid from torch_geometric.transforms import NormalizeFeatures dataset = Planetoid(root="data/Planetoid", name='CiteSeer', transform=NormalizeFeatures())
模型定义:
class GraphClassifier(torch.nn.Module): def __init__(self, dataset, hidden_dim): super(GraphClassifier, self).__init__() self.conv1 = GCNConv(dataset.num_features, hidden_dim) self.conv2 = GCNConv(hidden_dim, dataset.num_classes) def forward(self, data): x, edge_index = data.x, data.edge_index x = F.relu(self.conv1(x, edge_index)) x = F.relu(self.conv2(x, edge_index)) return F.log_softmax(x, dim=1)
可视化函数:
%matplotlib inline import matplotlib.pyplot as plt from sklearn.manifold import TSNE import torch # emb: (nNodes, hidden_dim) # node_type: (nNodes,). Entries are torch.int64 ranged from 0 to num_class - 1 def visualize(emb: torch.tensor, node_type: torch.tensor): z = TSNE(n_components=2).fit_transform(emb.detach().cpu().numpy()) plt.figure(figsize=(10,10)) plt.scatter(z[:, 0], z[:, 1], s=70, c=node_type, cmap="Set2") plt.show()
解决方案
1. 获取节点类型(node_type)
CiteSeer数据集的节点真实标签直接存储在数据集的data.y中,完全符合visualize函数对node_type的要求:
data = dataset[0] node_type = data.y
2. 获取最后一层节点嵌入(emb)
当前模型的forward方法仅返回分类对数概率,我们需要提取conv2层ReLU激活后的输出作为节点嵌入,有两种实现方式:
方式一:修改模型forward方法(推荐)
调整模型定义,让forward同时返回分类结果和节点嵌入,训练和可视化场景都能兼容:
import torch.nn.functional as F from torch_geometric.nn import GCNConv class GraphClassifier(torch.nn.Module): def __init__(self, dataset, hidden_dim): super(GraphClassifier, self).__init__() self.conv1 = GCNConv(dataset.num_features, hidden_dim) self.conv2 = GCNConv(hidden_dim, dataset.num_classes) def forward(self, data): x, edge_index = data.x, data.edge_index x = F.relu(self.conv1(x, edge_index)) emb = F.relu(self.conv2(x, edge_index)) # 这就是最后一层节点嵌入 logits = F.log_softmax(emb, dim=1) return logits, emb # 同时返回分类结果与嵌入
可视化时调用模型获取嵌入:
model.eval() # 切换到评估模式 with torch.no_grad(): # 关闭梯度计算节省资源 _, emb = model(data)
方式二:手动执行前向传播(无需修改模型)
如果不想改动现有模型,可以直接调用模型的层计算嵌入:
model.eval() with torch.no_grad(): x = data.x x = F.relu(model.conv1(x, data.edge_index)) emb = F.relu(model.conv2(x, data.edge_index)) # 得到最后一层嵌入
3. 执行可视化
获取emb和node_type后,直接传入可视化函数即可:
visualize(emb, node_type)
内容的提问来源于stack exchange,提问作者Peyman
相关产品推荐
相关产品推荐

