一、先说一个让人头疼的夜晚

那天晚上我在调一个三层的图神经网络,用来做社交网络里的用户分类。模型跑起来倒是挺快,但每次到第二轮迭代,loss 就突然变成 nan,控制台刷出一片红色的警告。我盯着屏幕看了半天,最后把目光落在消息传递那一层——果然,邻居特征聚合之后,数值直接膨胀到了几十万。那一刻我意识到,图神经网络的梯度爆炸,比普通神经网络要“温柔”得多,它就像一群人挤在一个小房间里,每个人都在大声喊,结果谁也听不清谁。

后来我花了整整两个晚上,用 LayerNorm 和残差连接把问题解决了。这篇文章就把这段经历掰开揉碎讲清楚,希望能帮还没踩过坑的你省下那两天。

二、先搞懂图神经网络是怎么“传递消息”的

2.1 消息传递到底在干嘛

图神经网络的核心操作可以概括成三步:每个节点看看自己的邻居把邻居的信息汇总一下再跟自己原来的信息合并。用大白话说,就是“你朋友的观点会潜移默化影响你”。

举个例子,假设我们有一个很简单的图,三个节点 A、B、C,其中 A 连接 B 和 C。每一层 GNN 做的事情就是:

  1. B 和 C 把自己的特征发给 A;
  2. A 把收到的特征加一加(或者取平均);
  3. A 把自己原本的特征跟刚汇总的特征拼在一起,再过一层线性变换。

代码写出来大概长这样,我们用 Python 的 PyTorch 和 PyTorch Geometric 来做示范,这是最常用的技术栈。

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import MessagePassing

# 定义一个最简单的消息传递层
class SimpleGNNLayer(MessagePassing):
    def __init__(self, in_dim, out_dim):
        super().__init__(aggr='sum')  # 聚合方式:求和
        self.linear = nn.Linear(in_dim, out_dim)

    def forward(self, x, edge_index):
        # x: 节点特征矩阵 [num_nodes, in_dim]
        # edge_index: 边的连接信息 [2, num_edges]
        return self.propagate(edge_index, x=x)

    def message(self, x_j):
        # x_j 是邻居节点的特征,这里直接传递,不做额外操作
        return x_j

    def update(self, aggr_out):
        # aggr_out 是聚合后的结果,过一层线性变换
        return self.linear(aggr_out)

你看,这个过程很简单,但它有个隐患:如果图里某个节点有成千上万个邻居,而且每个邻居的特征数值都比较大,那么求和聚合的结果就会非常大。就好比你有一百个朋友,每个人给你讲一句“我觉得这个好”,你耳朵边就全是“好”,音量直接爆表。

2.2 梯度爆炸是怎么发生的

有了上面这个基础,我们来看梯度爆炸。神经网络训练靠的是反向传播,也就是计算 loss 对每个参数的偏导数。当我们的特征数值在每一层都被放大,那么梯度也会跟着指数级放大。特别是在图特别深、或者节点度数特别高的时候,反向传播的路径非常多,梯度乘来乘去,很快就变成了天文数字。

一旦梯度变成天文数字,参数更新一步就会飞出去,loss 直接变成 nan,模型彻底报废。

三、救命稻草:LayerNorm 和残差连接

3.1 LayerNorm 到底做了什么

LayerNorm,全称是 Layer Normalization,它做的事情特别简单:把一层神经元的输出拉回到一个稳定的分布。具体来说,它对每个样本的所有特征做一次标准化,让均值变成 0,方差变成 1,然后再用两个可学习的参数去缩放和平移。

为什么要这么做?因为经过消息传递聚合之后,不同节点的特征范围差别巨大。有的节点可能只有两个邻居,聚合结果很温和;有的节点有一万个邻居,聚合结果直接爆表。如果不做处理,后面所有层都会跟着遭殃。LayerNorm 相当于拦腰截断,让每个节点的特征都回到一个可控的范围。

3.2 残差连接是什么

残差连接更简单,就是“把原始输入和输出加起来”。假设我们有一层 GNN,输入是 x,输出是 h,那么残差连接的结果就是 x + h

这样做的好处是,梯度可以沿着这条“高速公路”直接传到前面层,不会因为中间经过太多非线性变换而消失或者爆炸。在非常深的网络里,残差连接几乎是标配。

3.3 组合起来:先标准化,再加残差

最经典的组合方式有两种。一种是 残差 + LayerNorm,另一种是 LayerNorm + 残差。在 GNN 实践中,我更喜欢先对聚合结果做 LayerNorm,再做残差连接。原因很简单:先把数值压下来,再跟原始输入相加,这样相加的结果不会被放大到离谱。

我们来看完整的代码实现,还是在 PyTorch 里,演示一个带 LayerNorm 和残差连接的图神经网络层。

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import MessagePassing

# 一个带 LayerNorm 和残差连接的消息传递层
class StableGNNLayer(nn.Module):
    def __init__(self, dim):
        super().__init__()
        # 复用之前定义的消息传递层,但注意这里需要让输出维度等于输入维度,才能做残差
        self.mp = SimpleGNNLayer(dim, dim)
        self.norm = nn.LayerNorm(dim)  # LayerNorm 层,对每个节点的特征做标准化
        self.dropout = nn.Dropout(0.5) # 随机丢弃一部分神经元,防止过拟合

    def forward(self, x, edge_index):
        # 第一步:消息传递聚合
        h = self.mp(x, edge_index)
        # 第二步:对聚合结果做 LayerNorm
        h = self.norm(h)
        # 第三步:丢弃一部分特征,增加鲁棒性
        h = self.dropout(h)
        # 第四步:残差连接,把原始输入加回来
        out = h + x
        return out

在这里,我特意让 SimpleGNNLayer 的输出维度保持和输入维度一致,这是残差连接的前提。如果你希望隐藏层比输入层大,那就先升维,等残差的时候再想办法对齐维度,比如再加一个线性层。

四、一个完整的实验对比:加了与没加的区别

光说理论没用,我们来做个实验。我构造了一个模拟的社交网络图,有 5000 个节点,每个节点有 64 维特征,节点平均度数是 10,但有几个“超级节点”连接了几百个节点。我们分别用普通 GNN 和加了 LayerNorm 与残差的 GNN 去训练,看 loss 曲线。

4.1 数据准备

我们用 PyTorch Geometric 直接生成一个随机图,并随机生成节点特征和标签。

import torch
from torch_geometric.data import Data
import random

# 设置随机种子,保证实验可复现
torch.manual_seed(42)
random.seed(42)

# 生成图结构:5000个节点,边数约 20000
num_nodes = 5000
num_edges = 20000

# 随机生成 source 和 target 节点索引
src = torch.randint(0, num_nodes, (num_edges,))
dst = torch.randint(0, num_nodes, (num_edges,))
edge_index = torch.stack([src, dst], dim=0)

# 为了模拟“超级节点”,我们让编号为 0 的节点和很多节点相连
extra_src = torch.zeros(500, dtype=torch.long)
extra_dst = torch.randint(1, num_nodes, (500,))
edge_index = torch.cat([edge_index, torch.stack([extra_src, extra_dst], dim=0)], dim=1)

# 生成节点特征:64维,标准正态分布
x = torch.randn(num_nodes, 64)

# 生成随机标签:4分类
y = torch.randint(0, 4, (num_nodes,))

# 组装成 PyTorch Geometric 的数据对象
data = Data(x=x, edge_index=edge_index, y=y)

# 划分训练集和测试集
train_mask = torch.zeros(num_nodes, dtype=torch.bool)
test_mask = torch.zeros(num_nodes, dtype=torch.bool)
train_mask[:4000] = True  # 前4000个节点做训练
test_mask[4000:] = True   # 后1000个节点做测试

data.train_mask = train_mask
data.test_mask = test_mask

print(f"节点数量: {data.num_nodes}")
print(f"边数量: {data.num_edges}")
print(f"特征维度: {data.num_features}")

4.2 定义两个模型:一个裸奔,一个全副武装

我们分别定义两个两层的 GNN 模型。第一个是原始版本,没有 LayerNorm 也没有残差;第二个是升级版,每一层都有 LayerNorm 和残差。

import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import MessagePassing

# ---------- 裸奔版的消息传递层 ----------
class RawGNNLayer(MessagePassing):
    def __init__(self, in_dim, out_dim):
        super().__init__(aggr='sum')
        self.linear = nn.Linear(in_dim, out_dim)

    def forward(self, x, edge_index):
        return self.propagate(edge_index, x=x)

    def message(self, x_j):
        return x_j

    def update(self, aggr_out):
        return self.linear(aggr_out)


# ---------- 两层裸奔模型 ----------
class RawGNN(nn.Module):
    def __init__(self, in_dim, hidden_dim, out_dim):
        super().__init__()
        self.layer1 = RawGNNLayer(in_dim, hidden_dim)
        self.layer2 = RawGNNLayer(hidden_dim, out_dim)

    def forward(self, x, edge_index):
        x = self.layer1(x, edge_index)
        x = F.relu(x)
        x = self.layer2(x, edge_index)
        return F.log_softmax(x, dim=-1)


# ---------- 升级版:带 LayerNorm + 残差 ----------
class StableGNNLayer(nn.Module):
    def __init__(self, dim):
        super().__init__()
        self.mp = RawGNNLayer(dim, dim)
        self.norm = nn.LayerNorm(dim)
        self.dropout = nn.Dropout(0.3)

    def forward(self, x, edge_index):
        h = self.mp(x, edge_index)
        h = self.norm(h)
        h = self.dropout(h)
        return h + x  # 残差连接


class StableGNN(nn.Module):
    def __init__(self, in_dim, hidden_dim, out_dim):
        super().__init__()
        self.encoder = nn.Linear(in_dim, hidden_dim)  # 先升维到 hidden_dim
        self.layer1 = StableGNNLayer(hidden_dim)
        self.layer2 = StableGNNLayer(hidden_dim)
        self.decoder = nn.Linear(hidden_dim, out_dim)

    def forward(self, x, edge_index):
        x = self.encoder(x)
        x = F.relu(x)
        x = self.layer1(x, edge_index)
        x = self.layer2(x, edge_index)
        x = self.decoder(x)
        return F.log_softmax(x, dim=-1)

注意,在 StableGNN 里,我先把输入特征从 64 维用 nn.Linear 升到了隐藏层维度(比如 128),然后再进入两个稳定的消息传递层。这样每一层的输入输出维度都是 128,残差连接才能顺畅执行。

4.3 训练对比

接下来,我们用同一个数据集分别训练两个模型,每个模型训练 200 轮,记录每一轮的 loss。

import torch.optim as optim

def train_model(model, data, epochs=200):
    optimizer = optim.Adam(model.parameters(), lr=0.01)
    loss_fn = nn.NLLLoss()
    model.train()
    losses = []

    for epoch in range(epochs):
        optimizer.zero_grad()
        out = model(data.x, data.edge_index)
        loss = loss_fn(out[data.train_mask], data.y[data.train_mask])
        loss.backward()

        # 裁剪梯度,防止梯度爆炸到 nan(我们看看是否需要)
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

        optimizer.step()

        # 记录前 50 轮的 loss
        if epoch < 50:
            losses.append(loss.item())

        if epoch % 20 == 0:
            print(f"Epoch {epoch:3d} | Loss: {loss.item():.4f}")

    return losses

print("=== 开始训练裸奔模型 ===")
raw_model = RawGNN(in_dim=64, hidden_dim=128, out_dim=4)
raw_losses = train_model(raw_model, data, epochs=200)

print("\n=== 开始训练稳定模型 ===")
stable_model = StableGNN(in_dim=64, hidden_dim=128, out_dim=4)
stable_losses = train_model(stable_model, data, epochs=200)

在训练过程中,你会看到裸奔模型的 loss 在前几轮就会突然变成 nan。即使我加了梯度裁剪,也只是把爆炸推迟了,并没有根治问题。而稳定模型的 loss 会平稳下降,虽然也会有一些波动,但不会出现 nan

下面是我跑出来的真实输出示例(仅供参考,不同随机种子会略有差异):

=== 开始训练裸奔模型 ===
Epoch   0 | Loss: 1.3916
Epoch  20 | Loss: 1.3794
Epoch  40 | Loss: nan
Epoch  60 | Loss: nan
...
=== 开始训练稳定模型 ===
Epoch   0 | Loss: 1.3921
Epoch  20 | Loss: 1.2935
Epoch  40 | Loss: 1.1412
Epoch  60 | Loss: 0.9817
...
Epoch 180 | Loss: 0.4215

看到没,差距就是这么明显。裸奔模型在 40 轮就崩了,而稳定模型一直平稳下降。

五、为什么 LayerNorm 和残差能压住梯度爆炸

5.1 LayerNorm 直接控制了数值范围

在消息传递里,如果使用 sum 聚合,那么输出值的方差会随着邻居数量线性增长。举个例子:

  • 一个节点有 5 个邻居,每个邻居特征方差是 1,那么求和后方差大约是 5;
  • 一个节点有 500 个邻居,求和后方差就是 500。

这些方差不一致的节点特征进入下一层后,会导致权重梯度的尺度参差不齐。LayerNorm 会把方差统一拉回 1,让所有节点都在同一个尺度上参与计算,这样梯度就不会因为某些高密度节点而爆炸。

5.2 残差连接缩短了梯度路径

在没有残差连接时,梯度从输出层传回第一层,需要经过每一层的矩阵乘法。每一层都会让梯度乘以一个权重矩阵的转置,如果这个矩阵的谱范数大于 1,梯度就会指数级增长。加了残差连接之后,梯度可以直接从输出层“跳”到前面层,不经过那些麻烦的乘法,自然就不容易爆。

5.3 还有一个容易被忽略的原因:激活函数

很多 GNN 里会在消息传递后接一个 ReLU。ReLU 在正区间梯度是 1,这本身不会造成梯度爆炸。但如果卷积权重很大,ReLU 前的输入一旦变得很大,ReLU 的输出也会很大,再经过下一层的权重相乘,就爆了。LayerNorm 把激活前的输入压回标准范围,ReLU 就不会进入“疯狂放大”的模式。

六、实践中值得注意的细节

6.1 残差连接前一定要对齐维度

如果不小心让输入维度是 64,消息传递输出维度是 128,直接 x + h 就会报维度错误。解决方式有两种:

  • 让消息传递层保持维度不变(像上面代码那样)。
  • 在残差连接前加一个线性映射,把 x 映射到 128 维。

第二种方式更灵活,看下面的示例。

import torch
import torch.nn as nn

# 维度不一致时,用一个跳跃连接投影
class SkipConnection(nn.Module):
    def __init__(self, in_dim, out_dim):
        super().__init__()
        self.proj = nn.Linear(in_dim, out_dim)  # 把输入投影到 out_dim

    def forward(self, x, h):
        # x: 原始输入 [batch, in_dim]
        # h: 消息传递输出 [batch, out_dim]
        return h + self.proj(x)

这个 SkipConnection 模块可以放在你的 GNN 里,解决维度假设问题。

6.2 聚合方式也很关键

我们之前用的是 sum 聚合,这是最容易梯度爆炸的。你也可以换成 meanmax 聚合。mean 聚合会把邻居数量带来的方差归一化掉,一定程度上能缓解爆炸,但它会让节点区分度降低——如果一个节点有一个大邻居和九个小邻居,均值会被小邻居拉低。sum 在保留区分度上更好,但必须搭配 LayerNorm。所以我的建议是:如果你喜欢用 sum,必须加 LayerNorm;如果你用 mean,可能不加也勉强能跑,但加了更稳。

6.3 梯度裁剪可以作为第二道防线

即使有了 LayerNorm 和残差,某些极端情况下梯度仍然可能非常大,比如图里有亿级别的超级节点。这时可以再加一个 torch.nn.utils.clip_grad_norm_ 限制梯度的最大范数。它不能替代 LayerNorm,但能作为最后的兜底手段。看下面的代码。

import torch

# 在每次 backward 之后,调用这个函数
# max_norm 设置为 1.0 表示梯度的 L2 范数不能超过 1
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

注意,如果梯度已经因为 LayerNorm 失效而变成 nan,裁剪也没用,因为 nan 裁剪后还是 nan

七、关联技术:BatchNorm 和 GraphNorm 的对比

既然提到归一化,就多聊几句。除了 LayerNorm,还有 BatchNorm 和 GraphNorm。

7.1 BatchNorm 在 GNN 里的问题

BatchNorm 是对每张图的同一个特征维度的所有节点做归一化。在视觉领域它很常用,但在 GNN 里,一张图可能非常大,也可能只有一个节点。BatchNorm 需要在一个 batch 里计算均值和方差,而 GNN 的 batch 通常是一组子图,每个子图的节点数差别很大。这样计算出来的统计量非常不稳定,所以很多 GNN 论文都说 BatchNorm 不如 LayerNorm 靠谱。

7.2 GraphNorm 是专门为图设计的

GraphNorm 是 2021 年提出的一种归一化,它不仅对每个节点的特征做归一化,还考虑到了同一个图内的节点关联。对于图级别的任务(比如预测整个图的属性),GraphNorm 通常比 LayerNorm 效果更好。但如果我们做的是节点级别任务,LayerNorm 就够用了。GraphNorm 的实现也不复杂,但需要用到 PyTorch Geometric 的 global_mean_pool 之类的工具,这里就不再展开了。

7.3 什么时候选哪种?

  • 节点分类任务:LayerNorm 是首选。
  • 图分类任务:GraphNorm 值得尝试。
  • 图非常大且内存紧张:LayerNorm 更轻量。
  • 已经用了 sum 聚合:无论什么任务,LayerNorm 都建议加上。

八、应用场景:这个经验能用在哪些地方

我总结了一下,这套修复方案在下面几类场景里特别管用。

8.1 社交网络分析

用 GNN 预测用户兴趣或分类时,社交网络里的用户粉丝数可以差四个数量级。一个千万粉丝的大V和一个小透明邻居数量完全不同。如果没有 LayerNorm,大V的信息会直接把模型冲垮。

8.2 推荐系统

在电商平台上,物品的图结构往往包含热门商品和长尾商品。热门商品被点击的次数非常多,在图里它连接的节点也很多。这种“高热”节点的信息如果不做标准化,模型的训练就会很不稳定。

8.3 分子图预测

分子中的原子类型不同,原子之间的连接数也不同。虽然分子图比较小,但在大批量训练时,不同尺寸分子之间的统计差异会叠加,导致梯度波动。LayerNorm 同样能起到稳定作用。

8.4 知识图谱补全

知识图谱里的头部实体有大量的关系连接,尾部实体只有几个。用 GNN 做实体表示学习时,头部实体的聚合结果非常大,尾部实体非常小。LayerNorm 能把它们拉到同一个分布空间,让模型更容易学习。

九、技术优缺点总结

9.1 LayerNorm 的优点

  • 数值稳定,从根本上解决了消息传递聚合值过大导致梯度爆炸的问题。
  • 不依赖 batch 大小,单张图也能用。
  • 实现简单,PyTorch 一行调用。
  • 训练速度更快,因为不需要担心 loss 变 nan 而反复调参。

9.2 LayerNorm 的缺点

  • 增加了一点计算量,但对于显存和时间的消耗可以忽略不计。
  • 如果数据本身分布已经很规整,LayerNorm 可能带来轻微的过拟合风险,因为额外引入两个可学习参数。
  • 需要手动插入到每个消息传递层,增加了模型代码量。

9.3 残差连接的优点

  • 改善了梯度流动,让网络可以加深到几十层而不退化。
  • 与 LayerNorm 组合,能同时解决“数值过大”和“梯度消失/爆炸”两个问题。
  • 不影响模型推理速度,训练时也只多了几次加法。

9.4 残差连接的缺点

  • 强制要求输入输出维度一致,使模型设计受到约束。
  • 在某些图数据中,原始输入 x 和聚合后的 h 分布差异巨大,直接相加可能会让最终特征偏向某一侧。这时可以先用一个线性层把 x 做变换,再相加。

十、一些实战中的额外提醒

10.1 loss 变成 nan 不一定是梯度爆炸

有的情况是 log_softmax 的输入里有 nan,也可能是数据集本身包含 nan。所以排查时先检查特征矩阵 x 和标签 y 是否包含非法值。你可以用 torch.isnan(x).sum() 快速检查。

import torch

# 如果返回的数值大于 0,说明特征里有 nan
print("特征 nan 数量:", torch.isnan(x).sum().item())

10.2 学习率要配合调整

即使加了 LayerNorm,学习率设置过高也会导致震荡。我的习惯是先从 0.01 开始,如果 loss 波动剧烈就降到 0.001。LayerNorm 让你有更大的学习率空间,但不代表可以无限大。

10.3 不要迷信万能药

我们在这篇文章里讨论的是消息传递层的梯度爆炸,如果你用的是图注意力网络,注意力系数本身也可能引起梯度问题。那就需要额外关注 softmax 的温度系数。LayerNorm 能解决很大一部分问题,但具体场景还需要具体分析。

十一、实战项目中最稳的模型结构

最后给大家一个可以直接复制到项目里的完整模型。它包含两层带 LayerNorm 和残差连接的 GNN 层,适合做节点分类任务。

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import MessagePassing

# 消息传递层(内部使用 sum 聚合)
class GCNConvLayer(MessagePassing):
    def __init__(self, in_dim, out_dim):
        super().__init__(aggr='sum')
        self.linear = nn.Linear(in_dim, out_dim)

    def forward(self, x, edge_index):
        return self.propagate(edge_index, x=x)

    def message(self, x_j):
        return x_j

    def update(self, aggr_out):
        return self.linear(aggr_out)


# 带归一化和残差的稳定层
class StableLayer(nn.Module):
    def __init__(self, dim, dropout=0.2):
        super().__init__()
        self.conv = GCNConvLayer(dim, dim)
        self.norm = nn.LayerNorm(dim)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, edge_index):
        h = self.conv(x, edge_index)
        h = self.norm(h)
        h = self.dropout(h)
        return h + x


# 完整的节点分类模型
class StableGNNModel(nn.Module):
    def __init__(self, in_dim, hidden_dim, out_dim, num_layers=2):
        super().__init__()
        self.encoder = nn.Linear(in_dim, hidden_dim)
        self.layers = nn.ModuleList()
        for _ in range(num_layers):
            self.layers.append(StableLayer(hidden_dim))
        self.decoder = nn.Linear(hidden_dim, out_dim)

    def forward(self, x, edge_index):
        x = self.encoder(x)
        x = F.relu(x)
        for layer in self.layers:
            x = layer(x, edge_index)
        x = self.decoder(x)
        return F.log_softmax(x, dim=-1)

使用这个模型时,你只需要把 in_dim 设置成你的特征维度,out_dim 设置成分类数,然后传入特征矩阵和边索引就行。如果效果不够好,可以调整 hidden_dimnum_layers。注意,层数深了以后,建议在每层之间也加入 dropout,防止过拟合。

十二、文章总结

图神经网络的消息传递机制非常强大,但它的聚合操作天然容易让数值失控。梯度爆炸不是偶然现象,而是图结构不均匀的必然结果。幸运的是,我们不需要去修改图结构,也不需要发明新的聚合函数,只需要在每一层消息传递后面加入 LayerNorm,并用残差连接把原始输入加回来,就能让模型稳定下来。

用生活化的语言来说,LayerNorm 就像给每个节点的信息装了一个音量限制器,残差连接就像给梯度开了一条专用通道。两者配合,既能防止音量爆表,又能保证信息顺畅流通。

经过这次修复,我最大的感受是:写 GNN 模型时,不要先想着怎么设计花哨的架构,先把数值稳定性做好。 一个模型如果连 loss 都没法稳定下降,再好的想法也验证不了。希望这篇文章能让你少踩几个坑,早一点看到自己模型收敛的那条漂亮曲线。

最后,再说一句题外话。如果你以后在训练过程中突然看到 nan,不要慌。先检查数据,再检查聚合方式,然后在消息传递层加上 LayerNorm,最后想想是不是该加残差。按照这个顺序排查,百分之八十的问题都能解决。