解码器 - 代码实现及测试
1 课程概览
本课讲解 Transformer 解码器的代码实现及测试。解码器由多个解码器层堆叠而成(默认 6 个),最后加一层规范化层。前向传播接收 X(解码器输入)、encoder_output(编码器输出)、source_mask(填充掩码)、target_mask(目标掩码)四个参数。
2 核心概念与定义
- Decoder 类:Transformer 解码器,由多个解码器层堆叠组成。
- 解码器层堆叠:默认 6 个解码器层。
- 规范化层:最后一层,让输出更稳定。
- source_mask:填充掩码,用于编码器-解码器注意力。
- target_mask:目标掩码,用于掩码多头自注意力。
3 模型与算法详解
解码器结构
X → [解码器层 1] → [解码器层 2] → ... → [解码器层 N] → LayerNorm → output
↑
encoder_output
source_mask, target_mask
前向传播参数
| 参数 | 说明 | 形状 |
|---|---|---|
| X | 解码器输入序列(词向量 + 位置编码) | [batch, seq_len, d_model] |
| encoder_output | 编码器输出序列 | [batch, seq_len, d_model] |
| source_mask | 填充掩码(用于编码器-解码器注意力) | [batch, 1, seq_len, seq_len] |
| target_mask | 目标掩码(用于掩码多头自注意力) | [batch, 1, seq_len, seq_len] |
掩码说明
| 掩码 | 作用 | 用于 |
|---|---|---|
| source_mask | 填充掩码(padding mask) | 编码器-解码器注意力 |
| target_mask | 目标掩码(sequence mask) | 掩码多头自注意力 |
4 数学原理与推导
解码器层堆叠
$$\text{X}{i} = \text{DecoderLayer}(\text{X}{i-1}, \text{encoder_output}, \text{source_mask}, \text{target_mask})$$
规范化层
$$\text{output} = \text{LayerNorm}(\text{X}_{N})$$
5 代码示例
import torch
import torch.nn as nn
import copy
class DecoderLayer(nn.Module):
"""解码器层"""
def __init__(self, size, self_attn, src_attn, feed_forward, dropout):
super(DecoderLayer, self).__init__()
self.size = size
self.self_attn = self_attn
self.src_attn = src_attn
self.feed_forward = feed_forward
self.sublayer = clone(SubLayerConnection, size, 3) # 三个子层
self.dropout = nn.Dropout(dropout)
def forward(self, x, memory, source_mask, target_mask):
m = memory
# 子层 1:掩码多头自注意力
x = self.sublayer[0](x, lambda x: self.self_attn(x, x, x, target_mask))
# 子层 2:编码器-解码器注意力
x = self.sublayer[1](x, lambda x: self.src_attn(x, m, m, source_mask))
# 子层 3:前馈全连接层
return self.sublayer[2](x, self.feed_forward)
class Decoder(nn.Module):
"""解码器:由多个解码器层堆叠组成"""
def __init__(self, layer, N):
"""
初始化函数
Args:
layer: 单个解码器层对象
N: 解码器层堆叠数量(默认 6)
"""
super(Decoder, self).__init__()
# 1. 复制 N 个解码器层
self.layers = clone(layer, N)
# 2. 规范化层
self.norm = nn.LayerNorm(layer.size)
def forward(self, x, encoder_output, source_mask, target_mask):
"""
前向传播
Args:
x: 解码器输入序列 [batch, seq_len, d_model]
encoder_output: 编码器输出序列 [batch, seq_len, d_model]
source_mask: 填充掩码(用于编码器-解码器注意力)
target_mask: 目标掩码(用于掩码多头自注意力)
Returns:
output: 解码器输出 [batch, seq_len, d_model]
"""
# 遍历每个解码器层
for layer in self.layers:
x = layer(x, encoder_output, source_mask, target_mask)
# 规范化层
return self.norm(x)
def clone(module, N):
"""复制 N 个模块"""
return nn.ModuleList([copy.deepcopy(module) for _ in range(N)])
class SubLayerConnection(nn.Module):
"""子层连接结构"""
def __init__(self, size, dropout):
super(SubLayerConnection, self).__init__()
self.norm = nn.LayerNorm(size)
self.dropout = nn.Dropout(dropout)
def forward(self, x, sublayer):
return x + self.dropout(sublayer(self.norm(x)))
# 测试
if __name__ == "__main__":
print("=== 解码器 - 代码实现及测试 ===")
d_model = 512
N = 6 # 解码器层数量
batch_size = 2
seq_len = 4
# 创建解码器层
self_attn = nn.MultiheadAttention(d_model, num_heads=8, batch_first=True)
src_attn = nn.MultiheadAttention(d_model, num_heads=8, batch_first=True)
feed_forward = nn.Sequential(
nn.Linear(d_model, 2048),
nn.ReLU(),
nn.Linear(2048, d_model)
)
decoder_layer = DecoderLayer(d_model, self_attn, src_attn, feed_forward, dropout=0.1)
# 创建解码器
decoder = Decoder(decoder_layer, N)
print(f"解码器:\n{decoder}")
# 创建输入
x = torch.randn(batch_size, seq_len, d_model) # 解码器输入
encoder_output = torch.randn(batch_size, seq_len, d_model) # 编码器输出
source_mask = torch.ones(batch_size, 1, seq_len, seq_len) # 填充掩码
target_mask = torch.tril(torch.ones(seq_len, seq_len)).unsqueeze(0).unsqueeze(0) # 目标掩码
target_mask = target_mask.expand(batch_size, 1, seq_len, seq_len)
# 前向传播
output = decoder(x, encoder_output, source_mask, target_mask)
print(f"\n输入形状: {x.shape}")
print(f"编码器输出形状: {encoder_output.shape}")
print(f"输出形状: {output.shape}") # [2, 4, 512]
代码说明
| 代码 | 说明 |
|---|---|
clone(layer, N) | 复制 N 个解码器层 |
nn.LayerNorm(layer.size) | 规范化层 |
self.layers | 解码器层列表 |
for layer in self.layers | 遍历每个解码器层 |
self.norm(x) | 规范化层 |
6 重难点与易错提醒
- ❗重点:解码器由多个解码器层堆叠组成(默认 6 个)。
- ❗重点:最后加一层规范化层,让输出更稳定。
- ❗重点:source_mask 用于编码器-解码器注意力,target_mask 用于掩码多头自注意力。
- ⚠️易错:打印解码器时不要加
layer,否则显示的是解码器层。
7 课堂问答精选
Q1:解码器由什么组成?
A:解码器由多个解码器层堆叠组成(默认 6 个),最后加一层规范化层。
Q2:前向传播接收哪些参数?
A:接收四个参数:x(解码器输入)、encoder_output(编码器输出)、source_mask(填充掩码)、target_mask(目标掩码)。
Q3:source_mask 和 target_mask 的区别?
A:source_mask 是填充掩码,用于编码器-解码器注意力;target_mask 是目标掩码,用于掩码多头自注意力。
8 本课小结
- 解码器由多个解码器层堆叠组成(默认 6 个)。
- 最后加一层规范化层,让输出更稳定。
- 前向传播接收 x、encoder_output、source_mask、target_mask 四个参数。
- source_mask 用于编码器-解码器注意力,target_mask 用于掩码多头自注意力。
9 延伸思考
- 如何构建输出部分?
- 如何搭建完整的 Transformer 架构?