从零手写大模型 · 推理篇 07 - Residual Connection(残差连接)

残差连接机制流程图
查看原图 ↗

上一期我们把多头注意力(Multi-Head Attention)讲完了,"猫、吃、鱼"这个3×4的矩阵,经过QKV变换、多头拆分、加权求和、拼接、乘WO之后,得到了一个新的3×4矩阵。

这一期开始,我们要把这个输出矩阵,和最开始的那个X,加在一起。

公式很简单:

output = X + MultiHeadAttention(X)

简单到很多人会跳过不想,但这一步恰恰是Transformer里最容易被低估的设计。这里的X,就是EP03做完Embedding、EP04加上Position之后的那个X——带着词的语义信息,也带着位置信息。


一、一个反直觉的问题

多头注意力的输出,本身已经是"经过X算出来的东西"了——Q、K、V全都来自X,WO也是对X衍生出的结果做的线性变换。

那问题来了:结果已经是X算出来的了,为什么还要在结果上再加一次原始的X?

这一步如果不做会怎样?答案不是"效果差一点",而是结构性的信息丢失,而且这个丢失会随着层数增加而加速恶化。


二、加权平均的代价:为什么会"丢信息"

回忆一下多头注意力在算什么:输出矩阵的第i行,是所有词的V向量按注意力权重做的加权求和。

以"猫吃鱼"为例,"猫"这一行的新表示,是"猫、吃、鱼"三个V向量按各自注意力权重混合出来的。"吃""鱼"这两行也是同样的操作,只是权重分布不同。

问题就出在"加权平均"这个动作本身:

只要是加权平均,就一定会让向量往"大家的均值"方向靠拢,每个词原本的独特性会被磨掉一部分。

单独看一层,这个磨损不明显。但Transformer是要堆很多层的。如果不加残差,这一层的输出矩阵直接作为下一层的输入,下一层又做一次加权平均……层数一多,"猫""吃""鱼"这三个原本语义、位置都不同的向量,会越来越像。

这个现象在数学上叫秩坍缩(rank collapse),而且退化速度是双指数级的,不是线性变的慢,而是塌缩得很快。

什么是"掉秩"?矩阵的秩,通俗理解就是矩阵里"有多少行是真正独立的信息,不能被其他行线性组合出来"。如果"猫"这一行和"吃"这一行越来越像,极端情况下"猫"这一行可以完全用"吃"这一行乘个系数表示出来,那这两行就不再是独立信息了,秩就要减一。秩往下掉,就意味着矩阵里能表达的独立信息量在变少。


三、残差连接怎么解决这个问题:Git的Diff思路

理解残差连接,可以类比Git的diff/commit机制:每一次commit,不是把整个文件重写一遍再存下来,而是只记录这次改动的patch,patch再叠加到base上。

残差连接做的是同一件事:

不让attention层去"重新生成"这个token的完整表示,而是只让它算一个"修正量/增量",再叠加在原始X上。

X就是那个base,里面带着Embedding的语义信息和Position的位置信息,会一直保留在主干上,不会因为某一层的算子操作被冲没。

即便某一层attention学得很差,算出来的修正量接近于0,残差这条路径也能保证"至少输出约等于输入",不会把之前层攒下来的信息全部推翻重来。这不是WO"学得不够详细"的问题——哪怕WO训练得再完美,只要"加权平均"这个数学操作在起作用,信息收敛就会发生,这是结构性的,跟训练质量无关。残差连接要解决的,正是这个结构性问题。


四、残差连接的第二个作用:给梯度开一条高速公路

除了保信息,残差连接还有一个同等重要的作用:防止梯度消失。

Transformer要堆很多层才有效果,GPT-2小模型就有12层。层数一深,反向传播时梯度要一层一层往回传,每经过一层矩阵乘法和非线性变换,梯度都可能被压缩一次,传到底层的时候可能已经小到几乎为0,模型学不动了。

残差连接的加法结构,相当于给梯度开了一条"直达高速公路":

对output = X + MultiHeadAttention(X)求导,X这一项的导数直接是1。梯度可以经过"+X"这条捷径直接传回去,不用非得穿过一整套矩阵乘法。这就是为什么Transformer能堆到十几层、几十层还能训练得动——没有残差连接,深层网络在训练层面基本训不动。

一边保底(防止信息被逐层磨平),一边开路(防止梯度逐层消失),这两个作用是同时成立的,不是二选一的关系。


五、残差连接的英文

残差连接的英文是 Residual Connection,也叫 Skip Connection(跳跃连接)。这个设计最早不是Transformer提出的,而是来自何恺明等人2015年的ResNet论文,用来解决当时图像分类深层CNN训练不动的问题。Transformer把这个思路借了过来,用在了每个子层(Attention、FFN)外面。


六、代码实现

残差连接本身其实就是一行加法,不需要什么复杂结构。这一期先不引入LayerNorm,就用最朴素的写法,把"X + 子层输出"这个动作单独跑一遍:

import torch

# 模拟"猫吃鱼":seq_len=3, d_model=8
X = torch.randn(3, 8)  # 原始输入:Embedding + Position 之后的X

# 假设这是多头注意力算出来的输出(EP06的结果),形状必须和X一致
attn_output = torch.randn(3, 8)

# 残差连接:就是一次矩阵加法
output = X + attn_output  # [3, 8],形状不变

print(X.shape, attn_output.shape, output.shape)
# torch.Size([3, 8]) torch.Size([3, 8]) torch.Size([3, 8])

如果想封装成模块,方便后面接入多头注意力和FFN,也可以简单包一层:

import torch.nn as nn

class ResidualConnection(nn.Module):
    def __init__(self, sublayer: nn.Module):
        super().__init__()
        self.sublayer = sublayer

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # x: [seq_len, d_model],例如 [3, 8] 对应"猫吃鱼"
        return x + self.sublayer(x)  # [3, 8],子层输出直接加在原始x上

这里的sublayer可以直接传入EP06写好的多头注意力模块,先不考虑LayerNorm。等下一期我们讲完LayerNorm之后,再回过头把这个模块升级成Pre-LN的完整写法。

需要提醒一句:目前模型权重还是随机初始化的,这一步矩阵加法本身没有任何语义上的意义,纯粹是结构上的演示。真正有意义的数值变化,要等到后面真实训练或者加载GPT-2预训练权重之后才会出现。


小结

  • 多头注意力的加权平均操作,会让不同词的向量表示逐渐趋同,层数越多越明显,极端情况下会导致矩阵秩坍缩;
  • 残差连接的本质是"只学增量,不重写全部",类似Git的diff/commit,保证原始信息不会被逐层冲刷掉;
  • 残差连接的第二个作用是给反向传播的梯度开一条捷径,防止深层网络梯度消失,训不动;
  • 这一期只讲残差本身,代码上就是一次加法,形状不变;LayerNorm和Pre-LN的写法放到下一期展开。

下一期:LayerNorm,讲完之后我们会把残差连接升级成GPT-2用的Pre-LN写法:x1 = x + Attention(LayerNorm(x))。

← 返回文章列表