完整的 Transformer 架构搭建(上)
1 课程概览
本课讲解完整的 Transformer 架构搭建(上)。Transformer 底层是编码器-解码器模型(EncoderDecoder 类),包含 source_embed(源输入嵌入层)、target_embed(目标输入嵌入层)、encoder(编码器)、decoder(解码器)、generator(输出生成器)五个主要组件。
2 核心概念与定义
- EncoderDecoder 类:Transformer 整体模型架构。
- source_embed:编码器的输入嵌入层(包含词嵌入和位置编码)。
- target_embed:解码器的输入嵌入层(包含词嵌入和位置编码)。
- encoder:编码器模块。
- decoder:解码器模块。
- generator:输出生成器(线性层 + LogSoftmax)。
3 模型与算法详解
Transformer 架构
源输入 → source_embed → encoder → encoder_output
↓
目标输入 → target_embed → decoder → output → generator → 概率分布
↑
encoder_output
主要组件
| 组件 | 说明 | 输入 | 输出 |
|---|---|---|---|
| source_embed | 源输入嵌入层 | [batch, seq_len] | [batch, seq_len, d_model] |
| target_embed | 目标输入嵌入层 | [batch, seq_len] | [batch, seq_len, d_model] |
| encoder | 编码器 | [batch, seq_len, d_model] | [batch, seq_len, d_model] |
| decoder | 解码器 | [batch, seq_len, d_model] | [batch, seq_len, d_model] |
| generator | 输出生成器 | [batch, seq_len, d_model] | [batch, seq_len, vocab_size] |
Seq2Seq 架构
Transformer 底层是 Seq2Seq 架构。
Seq2Seq 架构:
1. 编码器(Encoder)
2. 解码器(Decoder)
3. 中间语义张量 C
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):
"""
初始化函数
Args:
encoder: 编码器模块
decoder: 解码器模块
source_embed: 源输入嵌入层(包含词嵌入和位置编码)
target_embed: 目标输入嵌入层(包含词嵌入和位置编码)
generator: 输出生成器(线性层 + LogSoftmax)
"""
super(EncoderDecoder, self).__init__()
# 1. 编码器的输入嵌入层
self.source_embed = source_embed
# 2. 解码器的输入嵌入层
self.target_embed = target_embed
# 3. 编码器
self.encoder = encoder
# 4. 解码器
self.decoder = decoder
# 5. 输出生成器
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: 源序列掩码
target_mask: 目标序列掩码
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: 源序列掩码
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: 源序列掩码(填充掩码)
target_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)
# 测试
if __name__ == "__main__":
print("=== 完整的 Transformer 架构搭建(上)===")
d_model = 512
vocab_size = 1000
batch_size = 2
seq_len = 4
# 创建各组件(简化版)
source_embed = nn.Sequential(
nn.Embedding(vocab_size, d_model),
# PositionalEncoding(d_model)
)
target_embed = nn.Sequential(
nn.Embedding(vocab_size, d_model),
# PositionalEncoding(d_model)
)
# 编码器和解码器(简化版)
encoder = nn.TransformerEncoder(
nn.TransformerEncoderLayer(d_model=d_model, nhead=8, batch_first=True),
num_layers=6
)
decoder = nn.TransformerDecoder(
nn.TransformerDecoderLayer(d_model=d_model, nhead=8, batch_first=True),
num_layers=6
)
# 输出生成器
generator = nn.Linear(d_model, vocab_size)
# 创建 Transformer 模型
transformer = EncoderDecoder(encoder, decoder, source_embed, target_embed, generator)
print(f"Transformer 模型:\n{transformer}")
# 创建输入
source_x = torch.randint(0, vocab_size, (batch_size, seq_len)) # [2, 4]
target_y = torch.randint(0, vocab_size, (batch_size, seq_len)) # [2, 4]
source_mask = None # 简化
target_mask = None # 简化
# 前向传播
output = transformer(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]
代码说明
| 代码 | 说明 |
|---|---|
self.source_embed | 源输入嵌入层 |
self.target_embed | 目标输入嵌入层 |
self.encoder | 编码器 |
self.decoder | 解码器 |
self.generator | 输出生成器 |
self.encode(...) | 编码器前向传播 |
self.decode(...) | 解码器前向传播 |
6 重难点与易错提醒
- ❗重点:Transformer 底层是编码器-解码器模型。
- ❗重点:包含五个主要组件:source_embed、target_embed、encoder、decoder、generator。
- ❗重点:source_embed 和 target_embed 包含词嵌入和位置编码。
- 💡技巧:实际开发中不写 MyTransformer,写 EncoderDecoder 更专业。
7 课堂问答精选
Q1:Transformer 的主要组件有哪些?
A:五个主要组件:source_embed(源输入嵌入层)、target_embed(目标输入嵌入层)、encoder(编码器)、decoder(解码器)、generator(输出生成器)。
Q2:为什么类名用 EncoderDecoder 而不是 MyTransformer?
A:Transformer 底层是编码器-解码器模型,写 EncoderDecoder 更专业,体现了 Seq2Seq 架构。
Q3:source_embed 和 target_embed 包含什么?
A:包含词嵌入层和位置编码层。
8 本课小结
- Transformer 底层是编码器-解码器模型(EncoderDecoder 类)。
- 包含五个主要组件:source_embed、target_embed、encoder、decoder、generator。
- source_embed 和 target_embed 包含词嵌入和位置编码。
- 前向传播:编码器 → 解码器 → 输出生成器。
9 延伸思考
- 如何实现编码器和解码器的前向传播?
- 如何测试完整的 Transformer 架构?