完整的 Transformer 架构搭建(下)
1 课程概览
本课讲解完整的 Transformer 架构搭建(下)。实现编码器和解码器的前向传播方法。编码器接收 source_x 和 source_mask,返回 encoder_output;解码器接收 target_y、encoder_output、source_mask、target_mask,返回解码器输出。source_mask 是填充掩码(padding mask),target_mask 是序列掩码(sequence mask)。
2 核心概念与定义
- encode 方法:编码器前向传播,返回 encoder_output。
- decode 方法:解码器前向传播,返回解码器输出。
- source_mask:填充掩码(padding mask),防止填充的 pad 值影响注意力计算结果。
- target_mask:序列掩码(sequence mask),防止未来的信息被提前利用(防止偷看未来的词)。
3 模型与算法详解
encode 方法
source_x [batch, seq_len] → source_embed → [batch, seq_len, d_model] → encoder → encoder_output [batch, seq_len, d_model]
decode 方法
target_y [batch, seq_len] → target_embed → [batch, seq_len, d_model] → decoder → output [batch, seq_len, d_model]
↑
encoder_output
source_mask, target_mask
掩码说明
| 掩码 | 作用 | 形状 | 用于 |
|---|---|---|---|
| source_mask | 填充掩码(padding mask) | [batch, 1, seq_len, seq_len] | 编码器-解码器注意力 |
| target_mask | 序列掩码(sequence mask) | [batch, 1, seq_len, seq_len] | 掩码多头自注意力 |
掩码的作用
source_mask(填充掩码):防止填充的 pad 值影响注意力计算结果。
target_mask(序列掩码):防止未来的信息被提前利用(防止偷看未来的词)。
target_mask(下三角矩阵):
1 0 0 0
1 1 0 0
1 1 1 0
1 1 1 1
4 数学原理与推导
编码器
$$\text{encoder_output} = \text{encoder}(\text{source_embed}(\text{source_x}), \text{source_mask})$$
解码器
$$\text{output} = \text{decoder}(\text{target_embed}(\text{target_y}), \text{encoder_output}, \text{source_mask}, \text{target_mask})$$
输出
$$\text{output} = \text{generator}(\text{output})$$
5 代码示例
import torch
import torch.nn as nn
class EncoderDecoder(nn.Module):
"""Transformer 整体模型架构"""
def __init__(self, encoder, decoder, source_embed, target_embed, generator):
super(EncoderDecoder, self).__init__()
self.source_embed = source_embed
self.target_embed = target_embed
self.encoder = encoder
self.decoder = decoder
self.generator = generator
def forward(self, source_x, target_y, source_mask, target_mask):
"""
前向传播
Args:
source_x: 源输入序列 [batch, seq_len]
target_y: 目标输入序列 [batch, seq_len]
source_mask: 填充掩码 [batch, 1, seq_len, seq_len]
target_mask: 序列掩码 [batch, 1, seq_len, seq_len]
Returns:
output: 概率分布 [batch, seq_len, vocab_size]
"""
# 1. 编码器前向传播
encoder_output = self.encode(source_x, source_mask)
# 2. 解码器前向传播
output = self.decode(target_y, encoder_output, source_mask, target_mask)
# 3. 输出生成器
return self.generator(output)
def encode(self, source_x, source_mask):
"""
编码器前向传播
Args:
source_x: 源输入序列 [batch, seq_len]
source_mask: 填充掩码(padding mask)
Returns:
encoder_output: 编码器输出 [batch, seq_len, d_model]
"""
# 1. 词嵌入 + 位置编码
embedded = self.source_embed(source_x)
# 2. 编码器处理
return self.encoder(embedded, source_mask)
def decode(self, target_y, encoder_output, source_mask, target_mask):
"""
解码器前向传播
Args:
target_y: 目标输入序列 [batch, seq_len]
encoder_output: 编码器输出 [batch, seq_len, d_model]
source_mask: 填充掩码(padding mask),用于编码器-解码器注意力
target_mask: 序列掩码(sequence mask),用于掩码多头自注意力
Returns:
output: 解码器输出 [batch, seq_len, d_model]
"""
# 1. 词嵌入 + 位置编码
embedded = self.target_embed(target_y)
# 2. 解码器处理
return self.decoder(embedded, encoder_output, source_mask, target_mask)
def make_model(source_vocab, target_vocab, N=6, d_model=512, d_ff=2048, num_heads=8, dropout=0.1):
"""
构建 Transformer 模型
Args:
source_vocab: 源词汇表大小
target_vocab: 目标词汇表大小
N: 编码器/解码器层数(默认 6)
d_model: 词向量维度(默认 512)
d_ff: 前馈全连接层维度(默认 2048)
num_heads: 多头注意力头数(默认 8)
dropout: 随机失活概率(默认 0.1)
Returns:
model: Transformer 模型
"""
import copy
# 1. 创建深拷贝函数
c = copy.deepcopy
# 2. 创建多头注意力层
attn = nn.MultiheadAttention(d_model, num_heads, dropout=dropout, batch_first=True)
# 3. 创建前馈全连接层
ff = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.ReLU(),
nn.Linear(d_ff, d_model)
)
# 4. 创建位置编码层
position = PositionalEncoding(d_model, dropout)
# 5. 创建编码器
encoder = nn.TransformerEncoder(
nn.TransformerEncoderLayer(d_model, num_heads, d_ff, dropout, batch_first=True),
N
)
# 6. 创建解码器
decoder = nn.TransformerDecoder(
nn.TransformerDecoderLayer(d_model, num_heads, d_ff, dropout, batch_first=True),
N
)
# 7. 创建源输入嵌入层
source_embed = nn.Sequential(nn.Embedding(source_vocab, d_model), c(position))
# 8. 创建目标输入嵌入层
target_embed = nn.Sequential(nn.Embedding(target_vocab, d_model), c(position))
# 9. 创建输出生成器
generator = nn.Linear(d_model, target_vocab)
# 10. 创建 Transformer 模型
model = EncoderDecoder(encoder, decoder, source_embed, target_embed, generator)
# 11. 参数初始化
for p in model.parameters():
if p.dim() > 1:
nn.init.xavier_uniform_(p)
return model
class PositionalEncoding(nn.Module):
"""位置编码层"""
def __init__(self, d_model, dropout, max_len=5000):
super(PositionalEncoding, self).__init__()
self.dropout = nn.Dropout(dropout)
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len).unsqueeze(1).float()
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-torch.log(torch.tensor(10000.0)) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0)
self.register_buffer('pe', pe)
def forward(self, x):
x = x + self.pe[:, :x.size(1)]
return self.dropout(x)
# 测试
if __name__ == "__main__":
print("=== 完整的 Transformer 架构搭建(下)===")
source_vocab = 1000
target_vocab = 1000
batch_size = 2
seq_len = 4
# 创建 Transformer 模型
model = make_model(source_vocab, target_vocab, N=6, d_model=512, num_heads=8)
print(f"Transformer 模型创建成功")
# 创建输入
source_x = torch.randint(0, source_vocab, (batch_size, seq_len)) # [2, 4]
target_y = torch.randint(0, target_vocab, (batch_size, seq_len)) # [2, 4]
source_mask = None # 简化
target_mask = None # 简化
# 前向传播
output = model(source_x, target_y, source_mask, target_mask)
print(f"\n源输入形状: {source_x.shape}")
print(f"目标输入形状: {target_y.shape}")
print(f"输出形状: {output.shape}") # [2, 4, 1000]
# 测试 encode 和 decode 方法
encoder_output = model.encode(source_x, source_mask)
print(f"\n编码器输出形状: {encoder_output.shape}") # [2, 4, 512]
decoder_output = model.decode(target_y, encoder_output, source_mask, target_mask)
print(f"解码器输出形状: {decoder_output.shape}") # [2, 4, 512]
代码说明
| 代码 | 说明 |
|---|---|
self.encode(source_x, source_mask) | 编码器前向传播 |
self.decode(target_y, encoder_output, source_mask, target_mask) | 解码器前向传播 |
self.source_embed(source_x) | 源输入词嵌入 + 位置编码 |
self.target_embed(target_y) | 目标输入词嵌入 + 位置编码 |
self.encoder(embedded, source_mask) | 编码器处理 |
self.decoder(embedded, encoder_output, source_mask, target_mask) | 解码器处理 |
make_model(...) | 构建 Transformer 模型 |
6 重难点与易错提醒
- ❗重点:encode 方法接收 source_x 和 source_mask。
- ❗重点:decode 方法接收 target_y、encoder_output、source_mask、target_mask。
- ❗重点:source_mask 是填充掩码(padding mask)。
- ❗重点:target_mask 是序列掩码(sequence mask)。
- 💡技巧:使用
make_model函数构建 Transformer 模型,参数初始化使用 Xavier。
7 课堂问答精选
Q1:encode 方法接收哪些参数?
A:接收 source_x(源输入序列)和 source_mask(填充掩码)。
Q2:decode 方法接收哪些参数?
A:接收 target_y(目标输入序列)、encoder_output(编码器输出)、source_mask(填充掩码)、target_mask(序列掩码)。
Q3:source_mask 和 target_mask 的区别?
A:source_mask 是填充掩码(padding mask),防止填充的 pad 值影响注意力计算结果;target_mask 是序列掩码(sequence mask),防止未来的信息被提前利用(防止偷看未来的词)。
8 本课小结
- encode 方法接收 source_x 和 source_mask,返回 encoder_output。
- decode 方法接收 target_y、encoder_output、source_mask、target_mask,返回解码器输出。
- source_mask 是填充掩码(padding mask)。
- target_mask 是序列掩码(sequence mask)。
- 使用
make_model函数构建 Transformer 模型。
9 延伸思考
- 如何测试完整的 Transformer 架构?
- 如何使用 Transformer 进行机器翻译?