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

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(图同构网络,理论上能表达更复杂的图结构)。
    • 加入全局信息:在节点特征中融入全局池化后的图级特征,或者使用全局注意力模块整合全局信息。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 19:48:14