别再死记硬背Transformer了!用PyTorch手把手实现一个简易翻译模型(附完整代码)
·
从零实现Transformer:用PyTorch构建英中翻译模型的实战指南
在深度学习领域,Transformer架构已经彻底改变了自然语言处理的格局。但对于许多学习者来说,阅读原始论文《Attention Is All You Need》后,面对复杂的架构图和数学公式,往往难以将理论转化为实践。本文将通过完整的代码实现,带你一步步构建一个精简但功能完整的英中翻译模型,让抽象的概念变得触手可及。
1. 环境准备与数据预处理
1.1 安装必要依赖
首先确保你的Python环境已安装最新版PyTorch。建议使用conda创建独立环境:
conda create -n transformer python=3.8
conda activate transformer
pip install torch torchtext spacy sentencepiece
1.2 准备双语数据集
我们将使用IWSLT 2017英中平行语料库的简化版本。为简化流程,已预处理好的数据包含约10,000个句子对:
import torchtext
from torchtext.data import Field, BucketIterator
# 定义字段处理器
SRC = Field(tokenize="spacy", tokenizer_language="en", init_token="<sos>", eos_token="<eos>", lower=True)
TRG = Field(tokenize="spacy", tokenizer_language="zh", init_token="<sos>", eos_token="<eos>", lower=True)
# 加载数据集
train_data, valid_data, test_data = torchtext.datasets.IWSLT.splits(
exts=('.en', '.zh'), fields=(SRC, TRG),
filter_pred=lambda x: len(vars(x)['src']) <= 50 and len(vars(x)['trg']) <= 50
)
# 构建词汇表
SRC.build_vocab(train_data, min_freq=2)
TRG.build_vocab(train_data, min_freq=2)
# 创建数据迭代器
BATCH_SIZE = 128
train_iterator, valid_iterator, test_iterator = BucketIterator.splits(
(train_data, valid_data, test_data), batch_size=BATCH_SIZE, device=device)
2. 模型核心组件实现
2.1 位置编码:捕捉序列顺序信息
Transformer摒弃了RNN的循环结构,需要通过位置编码注入序列顺序信息。以下是改进后的实现:
class PositionalEncoding(nn.Module):
def __init__(self, d_model, dropout=0.1, max_len=5000):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
position = torch.arange(max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
pe = torch.zeros(max_len, d_model)
pe[:, 0::2] = torch.sin(position * div_term) # 偶数位置使用sin
pe[:, 1::2] = torch.cos(position * div_term) # 奇数位置使用cos
pe = pe.unsqueeze(0).transpose(0, 1) # 形状: [max_len, 1, d_model]
self.register_buffer('pe', pe)
def forward(self, x):
"""
参数:
x: 输入张量,形状 [seq_len, batch_size, embedding_dim]
"""
x = x + self.pe[:x.size(0), :]
return self.dropout(x)
注意:位置编码的维度必须与词嵌入维度一致,这样才能直接相加。正弦和余弦函数的交替使用可以让模型更容易学习相对位置信息。
2.2 多头注意力机制实现
多头注意力是Transformer的核心创新,让我们分解实现:
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, n_head, dropout=0.1):
super().__init__()
assert d_model % n_head == 0, "d_model必须能被n_head整除"
self.d_k = d_model // n_head
self.n_head = n_head
self.w_q = nn.Linear(d_model, d_model)
self.w_k = nn.Linear(d_model, d_model)
self.w_v = nn.Linear(d_model, d_model)
self.w_o = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, q, k, v, mask=None):
batch_size = q.size(0)
# 线性投影并分头
q = self.w_q(q).view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2)
k = self.w_k(k).view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2)
v = self.w_v(v).view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2)
# 计算注意力分数
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
# 应用掩码(解码器自注意力使用)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 计算注意力权重
attn = F.softmax(scores, dim=-1)
attn = self.dropout(attn)
# 应用注意力权重到value上
output = torch.matmul(attn, v)
# 合并多头
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.n_head * self.d_k)
# 最终线性变换
return self.w_o(output)
关键点解析:
- 分头处理:将d_model维度的输入分割为n_head个较小的d_k维度子空间
- 缩放点积:注意力分数除以√d_k防止梯度消失
- 掩码机制:解码器中使用上三角掩码防止信息泄露
2.3 前馈网络与残差连接
每个注意力层后都跟随一个前馈网络和残差连接:
class PositionwiseFeedForward(nn.Module):
def __init__(self, d_model, d_ff, dropout=0.1):
super().__init__()
self.linear1 = nn.Linear(d_model, d_ff)
self.linear2 = nn.Linear(d_ff, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
return self.linear2(self.dropout(F.relu(self.linear1(x))))
class SublayerConnection(nn.Module):
"""残差连接后接层归一化"""
def __init__(self, size, dropout):
super().__init__()
self.norm = nn.LayerNorm(size)
self.dropout = nn.Dropout(dropout)
def forward(self, x, sublayer):
return x + self.dropout(sublayer(self.norm(x)))
3. 编码器与解码器架构
3.1 编码器层堆叠
编码器由N个相同层堆叠而成,每层包含:
- 多头自注意力机制
- 前馈神经网络
- 残差连接和层归一化
class EncoderLayer(nn.Module):
def __init__(self, d_model, n_head, d_ff, dropout):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, n_head, dropout)
self.feed_forward = PositionwiseFeedForward(d_model, d_ff, dropout)
self.sublayer = nn.ModuleList([SublayerConnection(d_model, dropout) for _ in range(2)])
def forward(self, x, mask):
x = self.sublayer[0](x, lambda x: self.self_attn(x, x, x, mask))
return self.sublayer[1](x, self.feed_forward)
class Encoder(nn.Module):
def __init__(self, vocab_size, d_model, n_head, d_ff, dropout, n_layers):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoding = PositionalEncoding(d_model, dropout)
self.layers = nn.ModuleList([EncoderLayer(d_model, n_head, d_ff, dropout) for _ in range(n_layers)])
self.norm = nn.LayerNorm(d_model)
def forward(self, src, src_mask):
x = self.pos_encoding(self.embedding(src))
for layer in self.layers:
x = layer(x, src_mask)
return self.norm(x)
3.2 解码器设计与实现
解码器不仅需要处理目标序列的自注意力,还要关注编码器输出:
class DecoderLayer(nn.Module):
def __init__(self, d_model, n_head, d_ff, dropout):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, n_head, dropout)
self.src_attn = MultiHeadAttention(d_model, n_head, dropout)
self.feed_forward = PositionwiseFeedForward(d_model, d_ff, dropout)
self.sublayer = nn.ModuleList([SublayerConnection(d_model, dropout) for _ in range(3)])
def forward(self, x, memory, src_mask, tgt_mask):
m = memory
x = self.sublayer[0](x, lambda x: self.self_attn(x, x, x, tgt_mask))
x = self.sublayer[1](x, lambda x: self.src_attn(x, m, m, src_mask))
return self.sublayer[2](x, self.feed_forward)
class Decoder(nn.Module):
def __init__(self, vocab_size, d_model, n_head, d_ff, dropout, n_layers):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoding = PositionalEncoding(d_model, dropout)
self.layers = nn.ModuleList([DecoderLayer(d_model, n_head, d_ff, dropout) for _ in range(n_layers)])
self.norm = nn.LayerNorm(d_model)
def forward(self, tgt, memory, src_mask, tgt_mask):
x = self.pos_encoding(self.embedding(tgt))
for layer in self.layers:
x = layer(x, memory, src_mask, tgt_mask)
return self.norm(x)
4. 模型训练与优化策略
4.1 损失函数与学习率调度
Transformer使用带标签平滑的交叉熵损失和动态学习率:
class LabelSmoothing(nn.Module):
def __init__(self, size, padding_idx, smoothing=0.1):
super().__init__()
self.criterion = nn.KLDivLoss(reduction='sum')
self.padding_idx = padding_idx
self.confidence = 1.0 - smoothing
self.smoothing = smoothing
self.size = size
def forward(self, x, target):
x = x.log_softmax(dim=-1)
true_dist = torch.zeros_like(x)
true_dist.fill_(self.smoothing / (self.size - 2))
true_dist.scatter_(1, target.unsqueeze(1), self.confidence)
true_dist[:, self.padding_idx] = 0
mask = torch.nonzero(target == self.padding_idx)
if mask.size(0) > 0:
true_dist.index_fill_(0, mask.squeeze(), 0.0)
return self.criterion(x, true_dist)
class NoamOpt:
"优化器包装器,实现动态学习率"
def __init__(self, model_size, factor, warmup, optimizer):
self.optimizer = optimizer
self._step = 0
self.warmup = warmup
self.factor = factor
self.model_size = model_size
self._rate = 0
def step(self):
self._step += 1
rate = self.rate()
for p in self.optimizer.param_groups:
p['lr'] = rate
self._rate = rate
self.optimizer.step()
def rate(self, step=None):
if step is None:
step = self._step
return self.factor * (self.model_size ** (-0.5) * min(step ** (-0.5), step * self.warmup ** (-1.5)))
4.2 训练循环实现
完整的训练过程包含以下关键步骤:
def train_epoch(model, train_iter, optimizer, criterion, clip):
model.train()
total_loss = 0
for batch in train_iter:
src = batch.src.transpose(0, 1) # [batch, seq_len]
trg = batch.trg.transpose(0, 1)
optimizer.optimizer.zero_grad()
# 创建掩码
src_mask = (src != SRC.vocab.stoi['<pad>']).unsqueeze(-2)
tgt_mask = make_std_mask(trg, TRG.vocab.stoi['<pad>'])
output = model(src, trg[:, :-1], src_mask, tgt_mask[:, :-1, :-1])
loss = criterion(output.contiguous().view(-1, output.size(-1)),
trg[:, 1:].contiguous().view(-1))
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), clip)
optimizer.step()
total_loss += loss.item()
return total_loss / len(train_iter)
def make_std_mask(tgt, pad_idx):
"创建序列掩码和填充掩码"
tgt_mask = (tgt != pad_idx).unsqueeze(-2)
seq_mask = torch.ones(1, tgt.size(1), tgt.size(1)).tril().bool().to(device)
return tgt_mask & seq_mask
4.3 模型评估与推理
推理时使用beam search提高翻译质量:
def greedy_decode(model, src, src_mask, max_len, start_symbol):
memory = model.encode(src, src_mask)
ys = torch.ones(1, 1).fill_(start_symbol).long().to(device)
for i in range(max_len-1):
out = model.decode(memory, src_mask,
ys,
make_std_mask(ys, TRG.vocab.stoi['<pad>']))
prob = model.generator(out[:, -1])
_, next_word = torch.max(prob, dim=1)
next_word = next_word.item()
ys = torch.cat([ys, torch.ones(1, 1).long().fill_(next_word).to(device)], dim=1)
if next_word == TRG.vocab.stoi['<eos>']:
break
return ys
def evaluate(model, val_iter, criterion):
model.eval()
total_loss = 0
with torch.no_grad():
for batch in val_iter:
src = batch.src.transpose(0, 1)
trg = batch.trg.transpose(0, 1)
src_mask = (src != SRC.vocab.stoi['<pad>']).unsqueeze(-2)
tgt_mask = make_std_mask(trg, TRG.vocab.stoi['<pad>'])
output = model(src, trg[:, :-1], src_mask, tgt_mask[:, :-1, :-1])
loss = criterion(output.contiguous().view(-1, output.size(-1)),
trg[:, 1:].contiguous().view(-1))
total_loss += loss.item()
return total_loss / len(val_iter)
5. 实战技巧与性能优化
5.1 混合精度训练
使用AMP(自动混合精度)加速训练并减少显存占用:
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
def train_epoch_amp(model, train_iter, optimizer, criterion, clip):
model.train()
total_loss = 0
for batch in train_iter:
src = batch.src.transpose(0, 1)
trg = batch.trg.transpose(0, 1)
optimizer.optimizer.zero_grad()
src_mask = (src != SRC.vocab.stoi['<pad>']).unsqueeze(-2)
tgt_mask = make_std_mask(trg, TRG.vocab.stoi['<pad>'])
with autocast():
output = model(src, trg[:, :-1], src_mask, tgt_mask[:, :-1, :-1])
loss = criterion(output.contiguous().view(-1, output.size(-1)),
trg[:, 1:].contiguous().view(-1))
scaler.scale(loss).backward()
scaler.unscale_(optimizer.optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), clip)
scaler.step(optimizer.optimizer)
scaler.update()
total_loss += loss.item()
return total_loss / len(train_iter)
5.2 模型量化与部署
训练完成后,可以通过量化减小模型体积:
# 动态量化
quantized_model = torch.quantization.quantize_dynamic(
model, {nn.Linear}, dtype=torch.qint8
)
# 保存量化模型
torch.save(quantized_model.state_dict(), "transformer_quantized.pth")
5.3 常见问题排查
在实现过程中可能会遇到以下典型问题:
- 梯度爆炸:添加梯度裁剪(
clip_grad_norm_)和使用更小的学习率 - 过拟合:增加dropout率、使用标签平滑、添加更多训练数据
- 训练缓慢:
- 检查矩阵运算是否在GPU上执行
- 使用更大的batch size
- 启用混合精度训练
- BLEU分数低:
- 检查数据预处理是否正确
- 尝试更大的模型或更长时间的训练
- 调整beam search参数
以下是一个典型训练过程的超参数配置参考:
| 参数 | 推荐值 | 说明 |
|---|---|---|
| d_model | 512 | 模型隐藏层维度 |
| n_head | 8 | 注意力头数 |
| d_ff | 2048 | 前馈网络中间维度 |
| n_layers | 6 | 编码器/解码器层数 |
| dropout | 0.1 | 随机失活率 |
| batch_size | 128 | 批处理大小 |
| warmup_steps | 4000 | 学习率预热步数 |
| label_smoothing | 0.1 | 标签平滑系数 |
在实际项目中,根据具体任务需求调整这些参数。例如处理长文本时可能需要增加d_model,而资源受限环境可以减少n_layers。
更多推荐


所有评论(0)