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

如何在基于UPFD数据集的GCN图分类任务中获取最终节点嵌入?

如何获取GCN模型处理后的最终节点嵌入(基于UPFD假新闻检测数据集)

问题背景

基于UPFD假新闻检测图数据集构建图分类GCN模型时,需要提取模型处理后的最终节点嵌入用于后续项目,但尝试打印模型处理后的节点嵌入时,发现和输入前的原始嵌入一致,不清楚如何正确获取经过卷积层后的节点特征。

用户当前代码如下:

current_file = '.'

train_dataset = UPFD(current_file, 'politifact', 'spacy', 'train', ToUndirected())
val_dataset = UPFD(current_file, 'politifact', 'spacy', 'val', ToUndirected())
test_dataset = UPFD(current_file, 'politifact', 'spacy', 'test', ToUndirected())

train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=128, shuffle=False)
test_loader = DataLoader(test_dataset, batch_size=128, shuffle=False)

# before_training = train_dataset[0].x
# print('Feature vector(node embedding) of datapoint #0 (before gtn):\n\t', train_dataset[0].x)


class GraphTransformer(torch.nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels,
                 concat=False):
        super().__init__()
        self.concat = concat

        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv2 = TransformerConv(hidden_channels, hidden_channels)
        self.conv3 = TransformerConv(hidden_channels, hidden_channels)

        if self.concat:
            self.lin0 = Linear(in_channels, hidden_channels)
            self.lin1 = Linear(2 * hidden_channels, hidden_channels)

        self.lin2 = Linear(hidden_channels, out_channels)

    def forward(self, x, edge_index, batch):
        h = self.conv1(x, edge_index).relu()
        h = self.conv2(h, edge_index).relu()
        h = self.conv3(h, edge_index).relu()
        h = global_max_pool(h, batch)

        if self.concat:
            # Get the root node (tweet) features of each graph:
            root = (batch[1:] - batch[:-1]).nonzero(as_tuple=False).view(-1)
            root = torch.cat([root.new_zeros(1), root + 1], dim=0)
            news = x[root]

            news = self.lin0(news).relu()
            h = self.lin1(torch.cat([news, h], dim=-1)).relu()

        h = self.lin2(h)
        return h.log_softmax(dim=-1)


device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = GraphTransformer(train_dataset.num_features, 128, train_dataset.num_classes, concat=True).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=0.01)

原因分析

  1. 模型未保留节点级嵌入:当前forward方法中,经过三次卷积得到的节点特征h,被global_max_pool压缩成了图级嵌入(每个图对应一个向量),最终返回的是分类结果,没有保留每个节点的最终特征。
  2. 原始数据集不会被修改:PyTorch Geometric的所有图操作(如GCNConv)都会生成特征的副本,不会直接修改原始数据集中的x字段,所以直接查看train_dataset[0].x永远是原始输入特征。

解决方案

修改模型的forward方法,使其同时返回分类结果和节点级最终嵌入;若只需节点嵌入,也可单独返回。同时需要处理批量数据中不同图的节点划分,确保能对应到每个原始图的节点。

修改后的模型代码

class GraphTransformer(torch.nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels,
                 concat=False):
        super().__init__()
        self.concat = concat

        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv2 = TransformerConv(hidden_channels, hidden_channels)
        self.conv3 = TransformerConv(hidden_channels, hidden_channels)

        if self.concat:
            self.lin0 = Linear(in_channels, hidden_channels)
            self.lin1 = Linear(2 * hidden_channels, hidden_channels)

        self.lin2 = Linear(hidden_channels, out_channels)

    def forward(self, x, edge_index, batch):
        # 保留经过三次卷积后的节点嵌入(这就是最终的节点级特征)
        node_embeddings = self.conv1(x, edge_index).relu()
        node_embeddings = self.conv2(node_embeddings, edge_index).relu()
        node_embeddings = self.conv3(node_embeddings, edge_index).relu()
        
        # 图分类用的全局嵌入
        h = global_max_pool(node_embeddings, batch)

        if self.concat:
            root = (batch[1:] - batch[:-1]).nonzero(as_tuple=False).view(-1)
            root = torch.cat([root.new_zeros(1), root + 1], dim=0)
            news = x[root]

            news = self.lin0(news).relu()
            h = self.lin1(torch.cat([news, h], dim=-1)).relu()

        h = self.lin2(h)
        # 同时返回分类结果和节点嵌入
        return h.log_softmax(dim=-1), node_embeddings

获取单张图的节点嵌入示例

# 取训练集中的第一张图
data = train_dataset[0].to(device)
model.eval()
with torch.no_grad():
    pred, node_embeds = model(data.x, data.edge_index, data.batch)

# node_embeds就是这张图所有节点的最终嵌入
print("最终节点嵌入形状:", node_embeds.shape)
print("第一个节点的最终嵌入:\n", node_embeds[0])

获取批量数据的节点嵌入并拆分到对应图

如果需要从DataLoader中批量获取节点嵌入,可以通过batch向量拆分每个图的节点:

model.eval()
with torch.no_grad():
    for batch_data in train_loader:
        batch_data = batch_data.to(device)
        pred, batch_node_embeds = model(batch_data.x, batch_data.edge_index, batch_data.batch)
        
        # 按batch拆分每个图的节点嵌入
        num_graphs = batch_data.num_graphs
        node_embeds_per_graph = []
        for i in range(num_graphs):
            # 获取当前图的所有节点索引
            mask = (batch_data.batch == i)
            graph_node_embeds = batch_node_embeds[mask]
            node_embeds_per_graph.append(graph_node_embeds)
        
        # node_embeds_per_graph中每个元素对应一个图的节点嵌入
        print("批量中第一个图的节点嵌入形状:", node_embeds_per_graph[0].shape)
        break

内容的提问来源于stack exchange,提问作者Prerk

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 20:42:38