用PyTorch逐行复现Transformer:从Harvard NLP的注释代码到你的第一个翻译模型
·
用PyTorch逐行构建Transformer:从理论到德语-英语翻译实战
在自然语言处理领域,Transformer架构已经成为现代序列建模的核心支柱。本文将带您从零开始实现一个完整的Transformer模型,并应用于实际的德语-英语翻译任务。不同于简单的代码注释,我们将重点关注如何将理论转化为可运行的实践项目。
1. Transformer核心组件实现
1.1 多头注意力机制
Transformer最核心的创新就是其注意力机制。让我们先实现缩放点积注意力:
def scaled_dot_product_attention(Q, K, V, mask=None):
"""
Q: 查询矩阵 [batch_size, num_heads, seq_len, d_k]
K: 键矩阵 [batch_size, num_heads, seq_len, d_k]
V: 值矩阵 [batch_size, num_heads, seq_len, d_v]
mask: 可选掩码 [seq_len, seq_len]
"""
d_k = Q.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn_weights = F.softmax(scores, dim=-1)
output = torch.matmul(attn_weights, V)
return output, attn_weights
注意:缩放因子1/√d_k对稳定训练至关重要,防止点积值过大导致softmax梯度消失。
多头注意力的实现将上述操作并行化:
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8, dropout=0.1):
super().__init__()
assert d_model % num_heads == 0
self.d_k = d_model // num_heads
self.num_heads = num_heads
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.num_heads, self.d_k).transpose(1,2)
K = self.W_k(K).view(batch_size, -1, self.num_heads, self.d_k).transpose(1,2)
V = self.W_v(V).view(batch_size, -1, self.num_heads, self.d_k).transpose(1,2)
# 计算注意力
x, attn = scaled_dot_product_attention(Q, K, V, mask)
# 合并多头
x = x.transpose(1,2).contiguous().view(batch_size, -1, self.num_heads * self.d_k)
return self.W_o(x)
1.2 位置编码与嵌入层
由于Transformer没有递归结构,我们需要显式地注入位置信息:
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * -(math.log(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):
return x + self.pe[:, :x.size(1)]
下表比较了不同位置编码方法的优劣:
| 编码类型 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|
| 正弦编码 | 可处理任意长度序列 | 固定模式,不可学习 | 通用Transformer |
| 学习编码 | 可适应数据特性 | 受限于最大长度 | 固定长度任务 |
| 相对编码 | 捕捉相对位置关系 | 实现复杂 | 长序列任务 |
2. 构建完整模型架构
2.1 编码器层实现
每个编码器层包含自注意力机制和前馈网络:
class EncoderLayer(nn.Module):
def __init__(self, d_model, num_heads, d_ff=2048, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, num_heads, dropout)
self.feed_forward = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.ReLU(),
nn.Linear(d_ff, d_model)
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask):
attn_output = self.self_attn(x, x, x, mask)
x = self.norm1(x + self.dropout(attn_output))
ff_output = self.feed_forward(x)
return self.norm2(x + self.dropout(ff_output))
2.2 解码器层实现
解码器层额外包含编码器-解码器注意力:
class DecoderLayer(nn.Module):
def __init__(self, d_model, num_heads, d_ff=2048, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, num_heads, dropout)
self.cross_attn = MultiHeadAttention(d_model, num_heads, dropout)
self.feed_forward = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.ReLU(),
nn.Linear(d_ff, d_model)
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.norm3 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, memory, src_mask, tgt_mask):
# 自注意力
attn_output = self.self_attn(x, x, x, tgt_mask)
x = self.norm1(x + self.dropout(attn_output))
# 编码器-解码器注意力
cross_output = self.cross_attn(x, memory, memory, src_mask)
x = self.norm2(x + self.dropout(cross_output))
# 前馈网络
ff_output = self.feed_forward(x)
return self.norm3(x + self.dropout(ff_output))
3. 数据准备与预处理
3.1 IWSLT数据集加载
我们使用torchtext加载IWSLT德语-英语数据集:
from torchtext.data import Field, BucketIterator
from torchtext.datasets import IWSLT
SRC = Field(tokenize="spacy", tokenizer_language="de",
init_token="<sos>", eos_token="<eos>", lower=True)
TRG = Field(tokenize="spacy", tokenizer_language="en",
init_token="<sos>", eos_token="<eos>", lower=True)
train_data, valid_data, test_data = IWSLT.splits(
exts=('.de', '.en'), fields=(SRC, TRG),
filter_pred=lambda x: len(vars(x)['src']) <= 100 and len(vars(x)['trg']) <= 100
)
SRC.build_vocab(train_data, min_freq=2)
TRG.build_vocab(train_data, min_freq=2)
3.2 批处理与掩码生成
为高效训练,我们需要实现批处理和注意力掩码:
def create_mask(src, trg, pad_idx):
src_mask = (src != pad_idx).unsqueeze(1).unsqueeze(2)
if trg is not None:
trg_mask = (trg != pad_idx).unsqueeze(1).unsqueeze(2)
seq_len = trg.size(1)
nopeak_mask = (1 - torch.triu(torch.ones(1, seq_len, seq_len), diagonal=1)).bool()
trg_mask = trg_mask & nopeak_mask
else:
trg_mask = None
return src_mask, trg_mask
4. 模型训练与评估
4.1 训练循环实现
我们使用带学习率热启动的Adam优化器:
class TransformerTrainer:
def __init__(self, model, optimizer, criterion, device):
self.model = model
self.optimizer = optimizer
self.criterion = criterion
self.device = device
def train_epoch(self, iterator):
self.model.train()
epoch_loss = 0
for i, batch in enumerate(iterator):
src = batch.src.to(self.device)
trg = batch.trg.to(self.device)
self.optimizer.zero_grad()
output = self.model(src, trg[:,:-1])
output_dim = output.shape[-1]
output = output.contiguous().view(-1, output_dim)
trg = trg[:,1:].contiguous().view(-1)
loss = self.criterion(output, trg)
loss.backward()
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1)
self.optimizer.step()
epoch_loss += loss.item()
return epoch_loss / len(iterator)
4.2 评估与翻译生成
实现贪婪解码算法生成翻译结果:
def translate_sentence(model, sentence, src_field, trg_field, device, max_len=50):
model.eval()
tokens = [token.lower() for token in sentence]
tokens = [src_field.init_token] + tokens + [src_field.eos_token]
src_indexes = [src_field.vocab.stoi[token] for token in tokens]
src_tensor = torch.LongTensor(src_indexes).unsqueeze(0).to(device)
with torch.no_grad():
encoder_outputs = model.encoder(src_tensor)
trg_indexes = [trg_field.vocab.stoi[trg_field.init_token]]
for i in range(max_len):
trg_tensor = torch.LongTensor(trg_indexes).unsqueeze(0).to(device)
with torch.no_grad():
output = model.decoder(trg_tensor, encoder_outputs)
pred_token = output.argmax(2)[:,-1].item()
trg_indexes.append(pred_token)
if pred_token == trg_field.vocab.stoi[trg_field.eos_token]:
break
trg_tokens = [trg_field.vocab.itos[i] for i in trg_indexes]
return trg_tokens[1:]
5. 注意力可视化与分析
理解模型如何关注输入序列的关键部分:
def plot_attention(attention, source, target):
fig = plt.figure(figsize=(12, 8))
ax = fig.add_subplot(111)
cax = ax.matshow(attention, cmap='bone')
fig.colorbar(cax)
ax.set_xticklabels([''] + source, rotation=90)
ax.set_yticklabels([''] + target)
ax.xaxis.set_major_locator(ticker.MultipleLocator(1))
ax.yaxis.set_major_locator(ticker.MultipleLocator(1))
plt.show()
在实际项目中,我发现德语-英语翻译任务中,模型特别关注动词位置和名词性别标记。例如,在翻译"der Hund läuft schnell"(狗跑得快)时,模型会:
- 首先关注主语"der Hund"
- 然后聚焦动词"läuft"
- 最后处理副词"schnell"
这种注意力模式与人类翻译过程惊人地相似。
更多推荐


所有评论(0)