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

PyTorch Geometric中GCN回归任务MSE Loss出现NaN问题求助

解决GCN回归任务MSE Loss出现NaN的问题

先排查最明显的模型错误

你的模型存在维度不匹配的问题:

  • conv3的输出维度是data.num_node_features,但后续linear1硬编码了输入维度为104,除非你的节点特征刚好是104维,否则这会导致前向/反向传播时出现数值异常。

修复方案:

如果是直接用conv3的输出做线性层输入,修改线性层定义:

# 替换原linear1的代码
self.linear1 = torch.nn.Linear(data.num_node_features, 1)

如果是想融合中间层特征(比如conv1和conv3的输出),需要在forward中拼接特征,并对应修改线性层维度:

class GCN(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = GCNConv(data.num_node_features, 100)
        self.conv2 = GCNConv(100, 16)
        self.conv3 = GCNConv(16, data.num_node_features)
        # 拼接后维度是100 + data.num_node_features,替换104
        self.linear1 = torch.nn.Linear(100 + data.num_node_features, 1)

    def forward(self, data):
        x, edge_index = data.x, data.edge_index

        h1 = self.conv1(x, edge_index)
        h = F.relu(h1)
        h = F.dropout(h, training=self.training)
        h = self.conv2(h, edge_index)
        h = self.conv3(h, edge_index)
        # 拼接conv1和conv3的输出
        h = torch.cat([h1, h], dim=1)
        h = self.linear1(h)
        h = h.tanh()
        return h

解决梯度爆炸/数值不稳定问题

即使修复了维度问题,以下操作也能避免NaN出现:

  1. 降低学习率
    当前0.01的学习率对于GCN来说偏高,容易引发梯度爆炸,尝试降低到0.001:

    optimizer = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=5e-4)
    
  2. 归一化输入和标签
    如果输入特征data.x或标签data.y数值范围过大,MSE损失会急剧增大,反向传播时梯度溢出:

    # 归一化输入特征(按列标准化)
    data.x = (data.x - data.x.mean(dim=0)) / data.x.std(dim=0)
    # 归一化标签(回归任务常用)
    y_mean = data.y.mean()
    y_std = data.y.std()
    data.y = (data.y - y_mean) / y_std
    

    注意测试时需要用保存的y_mean和y_std还原真实预测值。

  3. 添加梯度裁剪
    限制梯度的最大范数,防止梯度爆炸:

    for epoch in range(3):
        optimizer.zero_grad()
        out = model(data)
        loss = F.mse_loss(out.squeeze(), data.y.squeeze())
        loss.backward()
        # 裁剪梯度,最大范数设为1.0
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
        optimizer.step()
        print(f'Epoch: {epoch}, Loss: {loss}')
    
  4. 检查数据本身是否含NaN
    确认输入数据和标签没有NaN值:

    print("data.x是否含NaN:", torch.isnan(data.x).any())
    print("data.y是否含NaN:", torch.isnan(data.y).any())
    

    如果有,需要先清理数据(比如用均值填充或删除异常样本)。

内容的提问来源于stack exchange,提问作者Md Tahmid Hasan Fuad

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 08:25:20