You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何提取并可视化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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 04:55:18