从零手写大模型 · 推理篇 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))。