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

基于GATv2Conv的模型损失停滞且节点预测值相同的问题排查

问题排查:GAT模型无法学习节点分类任务

任务与数据集概况

使用PyTorch构建几何深度学习模型,处理约5000个全连接图的节点分类任务,每个图仅含1个正节点(正确节点)。数据集已划分训练/验证/测试集并分批加载,训练加载器中一个批次示例:

DataBatch(x=[1404, 5], edge_index=[2, 14700], edge_attr=[2], y=[1], batch=[1404], ptr=[65])

节点特征x与标签y的形状及示例:

torch.Size([1406, 5])
tensor([[ 0.7833,  0.1309, -0.0708, -0.0496,  0.7143],
    [ 0.9170, -0.0228, -0.2538,  0.0542, -0.0476],
    [ 0.8326, -0.2361, -0.0492, -0.1727, -0.9048],
    ...,
    [ 0.9281,  0.0207, -0.1396,  0.0936,  0.1429],
    [ 0.9427,  0.8991, -0.1633,  0.0857, -1.0000],
    [ 0.8982,  0.0480, -0.4320,  0.0886, -0.2381]])
torch.Size([1406])
tensor([0., 0., 0.,  ..., 1., 0., 0.])

当前模型代码

import time

node_features = ['smooth_x', 'smooth_y', 'keeper', 'players_between', 'team']
edge_features = ['same_team', 'distance']

# Create the train, validation and test datasets
train_dataset, train_loader = create_train_batches(corner_graphs_train['Graph'], node_features, edge_features)
val_dataset, test_dataset, val_loader, test_loader = create_val_test_batches(corner_graphs_val['Graph'], corner_graphs_test['Graph'], node_features, edge_features)


class Net(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = GATv2Conv(train_dataset.num_features, 256, heads=8, dropout=0.3)

        self.conv2 = GATv2Conv(8 * 256, 256, heads=8, dropout=0.3)
        
        # Add dropout layer
        self.dropout = torch.nn.Dropout(p=0.3)

        self.conv3 = GATv2Conv(8* 256, 1)


    def forward(self, x, edge_index, batch):
        # Layer 1
        x = F.leaky_relu(self.conv1(x, edge_index))

        # Apply dropout
        x = self.dropout(x)

        # Layer 2
        x = F.leaky_relu(self.conv2(x, edge_index))

        # Apply dropout
        x = self.dropout(x)

        # Layer 3
        x = F.leaky_relu(self.conv3(x, edge_index))

        # From flat list of nodes to 64 lists of nodes belonging to the same graph based on batch
        x, mask = to_dense_batch(x, batch)
        x = x.squeeze(-1)
        return x


device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(device)
model = Net().to(device)
#loss_op = torch.nn.CrossEntropyLoss()
loss_op = torch.nn.BCEWithLogitsLoss()
#optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=5e-4)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=5e-4)


def train():
    model.train()

    total_loss = 0
    data_progress = 0
    for data in train_loader:
        #print('Data progress: {}/{}'.format(data_progress, len(train_loader)))
        optimizer.zero_grad()
        output = model(data.x.to(device), data.edge_index.to(device), data.batch.to(device))
        if data_progress == 1:
            print('Example prediction')
            print(list(output[0].float()))
        data_progress += 1
        truth, mask = to_dense_batch(data.y.to(device), data.batch.to(device))
        #print(output)
        #print(truth)
        loss = loss_op(output, truth)                 
        total_loss += loss.item() * data.num_graphs
        loss.backward()
        optimizer.step()
    return total_loss / len(train_loader.dataset)


@torch.no_grad()
def val(loader):
    model.eval()

    total_loss_val = 0
    ys, preds = [], []
    for data in loader:
        truth, mask = to_dense_batch(data.y.to(device), data.batch.to(device))
        ys.append(truth.float().cpu())
        out = model(data.x.to(device), data.edge_index.to(device), data.batch.to(device))
        preds.append((out > 0).float().cpu())

        loss = loss_op(out, truth)
        total_loss_val += loss.item() * data.num_graphs

    y, pred = torch.cat(ys, dim=0).numpy(), torch.cat(preds, dim=0).numpy()

    return f1_score(y, pred, average='samples'), total_loss_val / len(loader.dataset)

@torch.no_grad()
def test(loader):
    model.eval()

    ranks = []
    for data in loader:
        out = model(data.x.to(device), data.edge_index.to(device), data.batch.to(device))
        pred_probs = list(torch.softmax(out[0], dim=0).flatten().cpu().numpy())
        pred_ranks = pd.DataFrame({'probabilities': pred_probs, 'y': data.y})
        pred_ranks.sort_values('probabilities', ascending=False, inplace=True)
        pred_ranks['Rank'] = range(1, len(pred_ranks) + 1)
        rank_predicted = pred_ranks[pred_ranks['y'] == 1]['Rank'].values[0]

        ranks.append(rank_predicted)

    return np.mean(ranks)


times = []
for epoch in range(1, 200):
    start = time.time()
    loss = train()
    temp_loss_list.append(loss)
    print(f'Time: {time.time() - start:.4f}s')
    val_f1, loss_val = val(val_loader)
    print(f'Time: {time.time() - start:.4f}s')
    test_avg_rank = test(test_loader)
    print(f'Time: {time.time() - start:.4f}s')
    
    print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}, Val_f1: {val_f1:.4f}, '
        f'Val Loss: {loss_val:.4f}, Test Avg Rank: {test_avg_rank:.4f}')
    
    # If the loss is not decreasing two times in a row, break the loop
    if epoch > 1 and temp_loss_list[-1] > temp_loss_list[-2]:
        if epoch > 2 and temp_loss_list[-2] > temp_loss_list[-3]:
            if epoch > 3 and temp_loss_list[-3] > temp_loss_list[-4]:
                if epoch > 4 and temp_loss_list[-4] > temp_loss_list[-5]:
                    if epoch > 5 and temp_loss_list[-5] > temp_loss_list[-6]:
                        break

    times.append(time.time() - start)
print(f"Median time per epoch: {torch.tensor(times).median():.4f}s")

训练输出与问题表现

模型未学习到有效信息,多数情况下对图中所有节点输出近乎相同的预测值,验证集F1始终为0,测试集平均排名接近随机水平:

Epoch: 001, Loss: 0.3710, Val_f1: 0.0000, Val Loss: 0.2000, Test Avg Rank: 11.2676
[-0.0572, -0.0553, -0.0580, -0.0571, -0.0573, -0.0556, -0.0582, -0.0534, -0.0637, -0.0579, -0.0575, -0.0573, -0.0566, -0.0561, -0.0570, -0.0580, -0.0579, -0.0579, -0.0579, -0.0579, -0.0567, -0.0579, -0.0530]

Epoch: 002, Loss: 0.1981, Val_f1: 0.0000, Val Loss: 0.1995, Test Avg Rank: 11.5412
[-2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348, -2.1348]

Epoch: 003, Loss: 0.1978, Val_f1: 0.0000, Val Loss: 0.1990, Test Avg Rank: 11.9824
[-3.7584, -3.7584, -3.7535, -3.7191, -3.7683, -3.7191, -4.0502, -3.7632, -3.7683, -3.7683, -3.7587, -3.7588, -3.7621, -3.7621, -4.8226, -3.7588, -3.7191, -3.7683, -3.7585, -3.7587, -3.7587, -3.7191]

Epoch: 004, Loss: 0.1973, Val_f1: 0.0000, Val Loss: 0.1989, Test Avg Rank: 12.0324
[-2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.8775, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917, -2.0917]

Epoch: 005, Loss: 0.1971, Val_f1: 0.0000, Val Loss: 0.2002, Test Avg Rank: 11.8971
[-2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -2.6774, -3.2843, -2.6774, -2.6774, -2.6774, -3.3673, -2.6774, -2.6774, -2.6774]

Epoch: 006, Loss: 0.1971, Val_f1: 0.0000, Val Loss: 0.1990, Test Avg Rank: 12.1912
[-3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155, -3.0155]

Epoch: 007, Loss: 0.1961, Val_f1: 0.0000, Val Loss: 0.1986, Test Avg Rank: 12.2324
[-3.1592, -3.1592, -3.1585, -3.1592, -3.1521, -3.1521, -3.1585, -3.1585, -3.1592, -3.1521, -3.1592, -3.1592, -3.1592, -3.1592, -3.1592, -3.1592, -3.1585, -3.1592, -3.1589, -5.4095, -3.1592, -3.1592]

已尝试调整网络架构、学习率、特征组合、优化器等方案,均无改善。


排查方向与修复建议

1. 核心模型结构问题

  • 未使用边特征:数据集包含edge_attr(same_team、distance),但当前GAT层未传入该参数,丢失关键结构信息。需修改模型初始化与前向传播:
    # 初始化时指定edge_dim
    self.conv1 = GATv2Conv(train_dataset.num_features, 256, heads=8, dropout=0.3, edge_dim=2)
    self.conv2 = GATv2Conv(8 * 256, 256, heads=8, dropout=0.3, edge_dim=2)
    self.conv3 = GATv2Conv(8* 256, 1, edge_dim=2)
    
    # 前向传播时传入edge_attr
    x = F.leaky_relu(self.conv1(x, edge_index, edge_attr=edge_attr))
    
  • 输出层激活错误:BCEWithLogitsLoss需要原始logits作为输入,需移除最后一层的LeakyReLU:
    # 原代码
    x = F.leaky_relu(self.conv3(x, edge_index))
    # 修改为
    x = self.conv3(x, edge_index)
    
  • 模型容量过载:5维节点特征对应256*8的通道数过大,易导致梯度消失或过拟合。建议先缩小规模,比如改为64通道、4头注意力,同时可添加残差连接缓解梯度消失。

2. 损失与类别不平衡问题

  • 类别权重缺失:每个图仅1个正节点,属于极度不平衡任务,需给正样本设置权重:
    # 假设每个图有N个节点,计算正样本权重
    pos_weight = torch.tensor([(N-1)/1], device=device)
    loss_op = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight)
    
    或直接使用FocalLoss替代BCEWithLogitsLoss,进一步平衡类别影响。

3. 训练与测试流程问题

  • **梯度有效性检查
相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 11:55:44