子层结构搭建 - 代码实现
1 课程概览
本课实现子层连接结构(SubLayerConnection)。子层是组件的封装,包含残差连接(Add)和规范化层(Norm)。编码器有 2 个子层(多头注意力子层、前馈全连接子层),解码器有 3 个子层。子层设计思想:功能上实现 Sublayer(x) + x,即残差连接。通过传入不同的子层对象(多头注意力或前馈全连接),实现不同的子层功能。
2 核心概念与定义
- 子层(SubLayer):组件的封装,包含残差连接和规范化层。
- SubLayerConnection:子层连接结构。
- 残差连接(Add):
Sublayer(x) + x,防止梯度消失。 - 规范化层(Norm):Layer Normalization。
- 子层对象:可以是多头注意力或前馈全连接。
3 模型与算法详解
子层结构
子层 = 组件(多头注意力/前馈全连接)+ 残差连接(Add)+ 规范化层(Norm)
编码器子层
| 子层 | 组件 |
|---|---|
| 子层1 | 多头注意力 + Add & Norm |
| 子层2 | 前馈全连接 + Add & Norm |
解码器子层
| 子层 | 组件 |
|---|---|
| 子层1 | 掩码多头注意力 + Add & Norm |
| 子层2 | 多头注意力 + Add & Norm |
| 子层3 | 前馈全连接 + Add & Norm |
子层设计思想
功能上实现
Sublayer(x) + x,即残差连接。
- 编码器/解码器层由子层拼接起来
- 像堆积木一样
- 方便管理
SubLayerConnection 类
SubLayerConnection(nn.Module)
├── __init__(d_model, dropout)
│ ├── norm(规范化层)
│ └── dropout
└── forward(x, sublayer)
├── sublayer_output = sublayer(x)
├── residual = x + dropout(sublayer_output)
└── return norm(residual)
4 数学原理与推导
残差连接
$$\text{output} = \text{Sublayer}(x) + x$$
残差连接 + 规范化
$$\text{output} = \text{LayerNorm}(\text{Sublayer}(x) + x)$$
另一种实现(先规范化,再残差)
$$\text{output} = \text{Sublayer}(\text{LayerNorm}(x)) + x$$
5 代码示例
import torch
import torch.nn as nn
from dm06_layer_norm import LayerNorm
class SubLayerConnection(nn.Module):
"""子层连接结构"""
def __init__(self, d_model, dropout=0.1):
"""
初始化子层连接结构
:param d_model: 词向量维度
:param dropout: 随机失活概率
"""
super().__init__()
# 1. 规范化层
self.norm = LayerNorm(d_model)
# 2. dropout 层
self.dropout = nn.Dropout(dropout)
def forward(self, x, sublayer):
"""
前向传播
:param x: 输入张量
:param sublayer: 子层函数(多头注意力或前馈全连接)
:return: 残差连接 + 规范化后的结果
"""
# 1. 子层处理:sublayer(x)
# 2. 残差连接:x + dropout(sublayer(x))
# 3. 规范化:norm(x + dropout(sublayer(x)))
return self.norm(x + self.dropout(sublayer(x)))
# 测试代码
def dm07_test_sublayer_connection():
"""测试子层连接结构"""
from dm04_attention import MultiHeadAttention
from dm05_feedforward import FeedForward
d_model = 512
# 1. 创建子层连接结构
sublayer_conn = SubLayerConnection(d_model)
# 2. 创建输入数据
x = torch.randn(2, 4, d_model)
# 3. 定义子层函数(多头注意力)
mha = MultiHeadAttention(d_model, 8)
def multi_head_attn_sublayer(x):
Q = K = V = x
return mha(Q, K, V)
# 4. 通过子层连接处理
output = sublayer_conn(x, multi_head_attn_sublayer)
print(f"输入形状: {x.shape}")
print(f"输出形状: {output.shape}")
return output
if __name__ == "__main__":
dm07_test_sublayer_connection()
代码说明
| 代码 | 说明 |
|---|---|
LayerNorm(d_model) | 规范化层 |
nn.Dropout(dropout) | 随机失活 |
sublayer(x) | 子层处理(传入函数) |
x + self.dropout(sublayer(x)) | 残差连接 |
self.norm(...) | 规范化 |
6 重难点与易错提醒
- ❗重点:子层 = 组件 + 残差连接 + 规范化层。
- ❗重点:残差连接
Sublayer(x) + x防止梯度消失。 - ❗重点:通过传入不同的子层对象,实现不同的子层功能。
- ⚠️易错:
sublayer是一个函数,不是对象。 - 💡深入理解:子层像堆积木一样拼接成编码器/解码器层。
7 课堂问答精选
Q1:子层连接结构的作用是什么?
A:将组件(多头注意力或前馈全连接)封装成子层,包含残差连接和规范化层。功能上实现 Sublayer(x) + x,即残差连接,防止梯度消失。
Q2:编码器和解码器各有几个子层?
A:编码器有 2 个子层(多头注意力子层、前馈全连接子层),解码器有 3 个子层(掩码多头注意力子层、多头注意力子层、前馈全连接子层)。
Q3:子层如何实现不同的功能?
A:通过传入不同的子层对象(函数)。如果传入多头注意力,就是多头注意力子层;如果传入前馈全连接,就是前馈全连接子层。
Q4:残差连接的公式是什么?
A:output = LayerNorm(Sublayer(x) + x),即子层处理结果加上输入,再进行规范化。
8 本课小结
- 子层 = 组件 + 残差连接 + 规范化层。
- 编码器 2 个子层,解码器 3 个子层。
- 残差连接:
Sublayer(x) + x。 - 通过传入不同的子层对象实现不同功能。
- 子层像堆积木一样拼接成编码器/解码器层。
9 延伸思考
- 如何测试子层连接结构?
- 编码器层如何搭建?