基于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个正节点,属于极度不平衡任务,需给正样本设置权重:
或直接使用FocalLoss替代BCEWithLogitsLoss,进一步平衡类别影响。# 假设每个图有N个节点,计算正样本权重 pos_weight = torch.tensor([(N-1)/1], device=device) loss_op = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight)
3. 训练与测试流程问题
- **梯度有效性检查
相关产品推荐
相关产品推荐

