PyTorch Geometric GCN节点分类全输出0问题及远距离依赖疑问
问题排查与解决方案
一、模型仅输出0的核心问题与修复
1. 输出层与激活函数不匹配
你的GCN最后一层输出维度设为1,但使用了F.log_softmax(x, dim=1),这是错误的:
- 单输出维度的二分类任务,应该用
sigmoid输出正类概率; - 如果坚持用
log_softmax,需要把最后一层输出维度改为2(对应两个类别的log概率)。
当前单维度下,log_softmax在dim=1计算时,每个节点只有一个值,softmax结果恒为1,log后输出全0,导致argmax结果始终是0。
2. 损失函数与输出不兼容
F.nll_loss要求输入是[num_nodes, num_classes]的log概率矩阵,你将输出flatten后变成一维,匹配逻辑错误:
- 若用单输出维度,改用
BCEWithLogitsLoss(无需手动加sigmoid,数值稳定性更好); - 若改用2维输出,可继续使用
nll_loss。
3. 类别不平衡未处理
正样本仅占5-10%,模型会天然偏向预测多数类(0)。需给正样本设置更高的损失权重,比如在损失函数中传入weight参数。
4. 训练逻辑错误
你当前逐个graph单独训练200 epoch,前一个graph的训练参数会被后一个覆盖,模型无法学习全局模式。应使用PyTorch Geometric的DataLoader批量加载数据,在每个epoch遍历所有graph训练。
修改后的代码示例
import pandas as pd import torch import torch.nn.functional as F from torch_geometric.data import Data, DataLoader from torch_geometric.nn import GCNConv TOTAL_TRAINING_DATA_SIZE = 10 VALIDATION_FACTOR = 0.2 TRAINING_DATA_SIZE = int(TOTAL_TRAINING_DATA_SIZE * (1 - VALIDATION_FACTOR)) def read_graph_data(amount): graphs = [] for i in range(amount): with pd.ExcelFile(f'training_data/test_{i}_graph_matrices.xlsx') as graph_file: nodes = pd.read_excel(graph_file, 'nodes', usecols=[1, 2, 3, 4, 5]) edges = pd.read_excel(graph_file, 'edges', usecols=[1, 2]) node_classifications = pd.read_excel(graph_file, 'classifications', dtype={'violates': bool}, usecols=[1]) graphs.append([nodes, edges, node_classifications]) return graphs def create_dataset(graphs): dataset = [] for i in range(TOTAL_TRAINING_DATA_SIZE): nodes, edges, node_classifications = graphs[i] edge_index = torch.tensor(edges.values, dtype=torch.long) node_features = torch.tensor(nodes.values, dtype=torch.float) # 把标签转为float类型,适配BCEWithLogitsLoss node_classifications_ = torch.tensor(node_classifications.values, dtype=torch.float).flatten() data = Data(x=node_features, y=node_classifications_, edge_index=edge_index.t().contiguous()) data.validate(raise_on_error=True) dataset.append(data) return dataset class GCN(torch.nn.Module): def __init__(self): super().__init__() self.conv1 = GCNConv(5, 16) self.conv2 = GCNConv(16, 1) # 单输出对应二分类正类概率 def forward(self, data): x, edge_index = data.x, data.edge_index x = self.conv1(x, edge_index) x = F.relu(x) x = F.dropout(x, training=self.training, p=0.3) # 降低dropout比例,避免过拟合 x = self.conv2(x, edge_index) return x # 直接输出logits,交给BCEWithLogitsLoss处理 print('reading training data') graphs = read_graph_data(TOTAL_TRAINING_DATA_SIZE) dataset = create_dataset(graphs) # 拆分训练和验证集 train_dataset = dataset[:TRAINING_DATA_SIZE] val_dataset = dataset[TRAINING_DATA_SIZE:] train_loader = DataLoader(train_dataset, batch_size=2, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=2) print('starting training') device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = GCN().to(device) # 改用Adam优化器,学习率调整为0.001 optimizer = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=5e-4) # 处理类别不平衡:正样本权重设为9(因为负样本占90%,正样本占10%,权重反比) pos_weight = torch.tensor([9.0]).to(device) criterion = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight) model.train() for epoch in range(200): total_loss = 0 for data in train_loader: data = data.to(device) optimizer.zero_grad() out = model(data) # out是[num_nodes,1],flatten后和data.y([num_nodes])匹配 loss = criterion(out.flatten(), data.y) loss.backward() optimizer.step() total_loss += loss.item() print(f'Epoch {epoch+1}, Loss: {total_loss/len(train_loader):.4f}') print('starting eval') model.eval() all_correct = 0 eval_test_case_amount = 0 with torch.no_grad(): for data in val_loader: data = data.to(device) out = model(data) # 用sigmoid转为概率,大于0.5则预测为1 pred = (torch.sigmoid(out.flatten()) > 0.5).long() print(pred.sum()) print(data.y.sum()) comparison_correct = (pred == data.y.long()) correct = comparison_correct.sum() all_correct += correct eval_test_case_amount += len(pred) acc = int(all_correct) / int(eval_test_case_amount) print(f'Accuracy: {acc:.4f}')
二、GCN对长距离依赖的捕捉能力
- 标准GCN的感受野等于网络层数:每一层GCN会聚合1跳邻居的信息,k层GCN最多能捕捉k跳的依赖。你当前的2层模型最多只能捕捉2跳关联,无法覆盖5-10跳的依赖。
- 解决方案:
- 加深GCN层数:增加到5-10层,但需加入残差连接避免梯度消失,比如使用
GCNConv默认开启的add_self_loops=True,或者手动添加残差:x = x + self.conv(x, edge_index)。 - 更换模型结构:选择更适合长距离依赖的模型,比如GraphSAGE(通过采样聚合多跳邻居)、GAT(注意力机制可聚焦关键远距离节点)、GIN(图同构网络,理论上能表达更复杂的图结构)。
- 加入全局信息:在节点特征中融入全局池化后的图级特征,或者使用全局注意力模块整合全局信息。
- 加深GCN层数:增加到5-10层,但需加入残差连接避免梯度消失,比如使用
内容的提问来源于stack exchange,提问作者lukstru
相关产品推荐
相关产品推荐

