如何在基于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)
原因分析
- 模型未保留节点级嵌入:当前
forward方法中,经过三次卷积得到的节点特征h,被global_max_pool压缩成了图级嵌入(每个图对应一个向量),最终返回的是分类结果,没有保留每个节点的最终特征。 - 原始数据集不会被修改: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
相关产品推荐
相关产品推荐

