🎯 课程主题
SublayerConnection 代码实现——子层连接模块(残差连接 + 层归一化)。
📝 核心知识点
1. 模块作用
- 概念说明:将子层(注意力/FFN)与残差连接、层归一化组合。
- 关键细节:
- Transformer 编码器和解码器中有重复结构
- 每个子层都经过:层归一化 → 子层计算 → Dropout → 残差连接
- 编码器有 2 个子层,解码器有 3 个子层
- 通过复用此模块构建完整模型
2. 代码结构
- 概念说明:SublayerConnection(简称 SC)的组成。
- 关键细节:
- 包含 LayerNorm 层
- 包含 Dropout 层
- 接收一个子层(sublayer)作为参数
- 子层可以是多头注意力或 FFN
3. 执行顺序
- 概念说明:先层归一化,再子层计算,最后残差连接。
- 关键细节:
- 代码顺序:
LayerNorm → sublayer → Dropout → + x - 理论顺序(Post-Norm):
sublayer → + x → LayerNorm - 代码采用 Pre-Norm 方式:
LayerNorm(x) → sublayer → + x - 从上一层角度看,是先残差连接再层归一化
- 代码顺序:
4. Pre-Norm vs Post-Norm
- 概念说明:两种不同的归一化位置。
- 关键细节:
- Post-Norm(原论文):$\text{LayerNorm}(x + \text{sublayer}(x))$
- Pre-Norm(本代码):$x + \text{sublayer}(\text{LayerNorm}(x))$
- Pre-Norm 训练更稳定,梯度流更好
- 现代大模型多采用 Pre-Norm
🧮 核心公式与推导
Pre-Norm(本代码):
$$\text{output} = x + \text{Dropout}(\text{sublayer}(\text{LayerNorm}(x)))$$
Post-Norm(原论文):
$$\text{output} = \text{LayerNorm}(x + \text{Dropout}(\text{sublayer}(x)))$$
残差连接的作用:
$$\frac{\partial L}{\partial x_l} = \frac{\partial L}{\partial x_L} \left(1 + \frac{\partial}{\partial x_l} \sum_{i=l}^{L-1} F(x_i, W_i)\right)$$
残差连接保证梯度可以直接回传到浅层,缓解梯度消失。
🏗️ 模型架构与数据流向
在编码器中的应用:
💻 代码实战
import torch
import torch.nn as nn
import torch.nn.functional as F
class SublayerConnection(nn.Module):
"""
子层连接模块 = LayerNorm + Dropout + 残差连接
将子层(如多头注意力、FFN)与残差连接和层归一化组合
"""
def __init__(self, size, dropout):
"""
size: 特征维度 d_model
dropout: dropout比例
"""
super().__init__()
self.norm = LayerNorm(size) # 层归一化
self.dropout = nn.Dropout(dropout) # Dropout
def forward(self, x, sublayer):
"""
x: [batch, seq_len, d_model] 输入
sublayer: 子层函数 (多头注意力或FFN)
return: [batch, seq_len, d_model]
执行顺序 (Pre-Norm):
1. LayerNorm(x)
2. sublayer(LayerNorm(x))
3. Dropout(sublayer(...))
4. x + Dropout(...)
"""
# Pre-Norm: 先归一化, 再子层计算, 最后残差连接
return x + self.dropout(sublayer(self.norm(x)))
class LayerNorm(nn.Module):
"""层归一化"""
def __init__(self, features, eps=1e-6):
super().__init__()
self.a_2 = nn.Parameter(torch.ones(features))
self.b_2 = nn.Parameter(torch.zeros(features))
self.eps = eps
def forward(self, x):
mean = x.mean(-1, keepdim=True)
std = x.std(-1, keepdim=True)
return self.a_2 * (x - mean) / (std + self.eps) + self.b_2
# 使用示例: 构建单个编码器层
class EncoderLayer(nn.Module):
"""单个编码器层 = 2个子层"""
def __init__(self, size, self_attn, feed_forward, dropout):
super().__init__()
self.self_attn = self_attn
self.feed_forward = feed_forward
# 克隆2个SublayerConnection (编码器有2个子层)
self.sublayer = clones(SublayerConnection(size, dropout), 2)
self.size = size
def forward(self, x, mask):
# 子层1: 多头自注意力
x = self.sublayer[0](
x, lambda x: self.self_attn(x, x, x, mask)
)
# 子层2: 前馈神经网络
x = self.sublayer[1](x, self.feed_forward)
return x
def clones(module, N):
"""克隆模块N次"""
import copy
return nn.ModuleList([copy.deepcopy(module) for _ in range(N)])
# 测试
if __name__ == "__main__":
d_model = 512
batch_size = 2
seq_len = 10
# 创建SublayerConnection
sc = SublayerConnection(d_model, dropout=0.1)
# 定义一个简单的子层 (如FFN)
ffn = nn.Sequential(
nn.Linear(d_model, 2048),
nn.ReLU(),
nn.Dropout(0.1),
nn.Linear(2048, d_model)
)
x = torch.randn(batch_size, seq_len, d_model)
output = sc(x, ffn)
print(f"输入: {x.shape}")
print(f"输出: {output.shape}") # [2, 10, 512]
print(f"残差连接保证形状一致: {x.shape == output.shape}")
⚠️ 常见问题与避坑指南
- 代码采用 Pre-Norm(先归一化再子层计算),与原论文 Post-Norm 不同
- 残差连接要求输入输出形状一致(都是 $d_{model}$)
sublayer是函数/可调用对象,通过 lambda 传入- 编码器克隆 2 次 SC,解码器克隆 3 次 SC
💡 个人总结与延伸
SublayerConnection 是 Transformer 的关键组件,将残差连接和层归一化封装为可复用模块。Pre-Norm 方式训练更稳定,被现代大模型广泛采用。残差连接的思想来自 ResNet,解决了深层网络的梯度消失问题,是深度学习的重要突破。