Transformer 架构组件介绍
1 课程概览
本课介绍 Transformer 模型的组件组装。将编码器输入处理(词嵌入 + 位置编码)、解码器输入处理(词嵌入 + 位置编码)、编码器、解码器、输出生成器等组件组装到 EncoderDecoder 类中,形成完整的 Transformer 架构。
2 核心概念与定义
- EncoderDecoder 类:自定义的 Transformer 架构(编码解码模型)。
- nn.Sequential:处理链,将多个层按顺序组合(流水线)。
- source_embed:编码器输入处理链(词嵌入 + 位置编码)。
- target_embed:解码器输入处理链(词嵌入 + 位置编码)。
- 组件组装:将各个组件塞入 EncoderDecoder 类。
3 模型与算法详解
Transformer 组件
source_x → source_embed → encoder → encoder_output
↓
target_y → target_embed → decoder → output → generator → 概率分布
↑
encoder_output
组件说明
| 组件 | 说明 | 包含 |
|---|---|---|
| source_embed | 编码器输入处理链 | 词嵌入层 + 位置编码层 |
| target_embed | 解码器输入处理链 | 词嵌入层 + 位置编码层 |
| encoder | 编码器 | N 个编码器层 |
| decoder | 解码器 | N 个解码器层 |
| generator | 输出生成器 | 线性层 + LogSoftmax |
nn.Sequential 处理链
nn.Sequential 是处理链,先走 A 再走 B,跟流水线一样。
# 编码器输入处理链:先词嵌入,再位置编码
source_embed = nn.Sequential(source_embedding, source_position)
# 解码器输入处理链:先词嵌入,再位置编码
target_embed = nn.Sequential(target_embedding, target_position)
4 数学原理与推导
编码器输入处理
$$\text{embedded} = \text{PositionalEncoding}(\text{Embedding}(\text{source_x}))$$
解码器输入处理
$$\text{embedded} = \text{PositionalEncoding}(\text{Embedding}(\text{target_y}))$$
编码器
$$\text{encoder_output} = \text{encoder}(\text{embedded}, \text{source_mask})$$
解码器
$$\text{output} = \text{decoder}(\text{embedded}, \text{encoder_output}, \text{source_mask}, \text{target_mask})$$
输出
$$\text{output} = \text{generator}(\text{output})$$
5 代码示例
import torch
import torch.nn as nn
import copy
class EncoderDecoder(nn.Module):
"""自定义的 Transformer 架构(编码解码模型)"""
def __init__(self, encoder, decoder, source_embed, target_embed, generator):
"""
初始化函数
Args:
encoder: 编码器
decoder: 解码器
source_embed: 编码器输入处理链(词嵌入 + 位置编码)
target_embed: 解码器输入处理链(词嵌入 + 位置编码)
generator: 输出生成器
"""
super(EncoderDecoder, self).__init__()
self.encoder = encoder
self.decoder = decoder
self.source_embed = source_embed
self.target_embed = target_embed
self.generator = generator
def forward(self, source_x, target_y, source_mask, target_mask):
"""前向传播"""
encoder_output = self.encode(source_x, source_mask)
output = self.decode(target_y, encoder_output, source_mask, target_mask)
return self.generator(output)
def encode(self, source_x, source_mask):
"""编码器前向传播"""
return self.encoder(self.source_embed(source_x), source_mask)
def decode(self, target_y, encoder_output, source_mask, target_mask):
"""解码器前向传播"""
return self.decoder(self.target_embed(target_y), 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 模型"""
c = copy.deepcopy
# 1. 词嵌入层
source_embedding = nn.Embedding(source_vocab, d_model)
target_embedding = nn.Embedding(target_vocab, d_model)
# 2. 位置编码层
position = PositionalEncoding(d_model, dropout)
# 3. 编码器
encoder = nn.TransformerEncoder(
nn.TransformerEncoderLayer(d_model, num_heads, d_ff, dropout, batch_first=True),
N
)
# 4. 解码器
decoder = nn.TransformerDecoder(
nn.TransformerDecoderLayer(d_model, num_heads, d_ff, dropout, batch_first=True),
N
)
# 5. 编码器输入处理链(词嵌入 + 位置编码)
source_embed = nn.Sequential(c(source_embedding), c(position))
# 6. 解码器输入处理链(词嵌入 + 位置编码)
target_embed = nn.Sequential(c(target_embedding), c(position))
# 7. 输出生成器
generator = nn.Linear(d_model, target_vocab)
# 8. 创建 Transformer 模型
model = EncoderDecoder(encoder, decoder, source_embed, target_embed, generator)
# 9. 参数初始化
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
# 构建 Transformer 模型
my_transformer = make_model(source_vocab, target_vocab, N=6, d_model=512, num_heads=8)
print(f"Transformer 模型:\n{my_transformer}")
# 解释各组件
print(f"\n=== 组件解释 ===")
print(f"1. source_embed (nn.Sequential): 编码器输入处理链")
print(f" - 词嵌入层 (Embedding)")
print(f" - 位置编码层 (PositionalEncoding)")
print(f"\n2. target_embed (nn.Sequential): 解码器输入处理链")
print(f" - 词嵌入层 (Embedding)")
print(f" - 位置编码层 (PositionalEncoding)")
print(f"\n3. encoder: 编码器(N 个编码器层)")
print(f"\n4. decoder: 解码器(N 个解码器层)")
print(f"\n5. generator: 输出生成器(线性层 + LogSoftmax)")
代码说明
| 代码 | 说明 |
|---|---|
nn.Sequential(c(source_embedding), c(position)) | 编码器输入处理链 |
nn.Sequential(c(target_embedding), c(position)) | 解码器输入处理链 |
EncoderDecoder(encoder, decoder, source_embed, target_embed, generator) | Transformer 模型 |
nn.init.xavier_uniform_(p) | Xavier 参数初始化 |
6 重难点与易错提醒
- ❗重点:source_embed 和 target_embed 是处理链(nn.Sequential)。
- ❗重点:处理链包含词嵌入层和位置编码层。
- ❗重点:EncoderDecoder 是自定义的 Transformer 架构。
- 💡技巧:使用深拷贝确保组件参数不共享。
7 课堂问答精选
Q1:source_embed 包含什么?
A:source_embed 是处理链(nn.Sequential),包含词嵌入层和位置编码层。
Q2:nn.Sequential 的作用是什么?
A:nn.Sequential 是处理链,将多个层按顺序组合,先走 A 再走 B,跟流水线一样。
Q3:EncoderDecoder 类包含哪些组件?
A:包含五个组件:encoder、decoder、source_embed、target_embed、generator。
8 本课小结
- Transformer 模型由 EncoderDecoder 类封装。
- source_embed 和 target_embed 是处理链(nn.Sequential),包含词嵌入层和位置编码层。
- 五个主要组件:encoder、decoder、source_embed、target_embed、generator。
- 使用深拷贝确保组件参数不共享。
9 延伸思考
- 如何测试 Transformer 模型?
- 如何使用 Transformer 进行机器翻译?