从零实现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)

关键点解析:

  1. 分头处理:将d_model维度的输入分割为n_head个较小的d_k维度子空间
  2. 缩放点积:注意力分数除以√d_k防止梯度消失
  3. 掩码机制:解码器中使用上三角掩码防止信息泄露

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个相同层堆叠而成,每层包含:

  1. 多头自注意力机制
  2. 前馈神经网络
  3. 残差连接和层归一化
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 常见问题排查

在实现过程中可能会遇到以下典型问题:

  1. 梯度爆炸:添加梯度裁剪(clip_grad_norm_)和使用更小的学习率
  2. 过拟合:增加dropout率、使用标签平滑、添加更多训练数据
  3. 训练缓慢
    • 检查矩阵运算是否在GPU上执行
    • 使用更大的batch size
    • 启用混合精度训练
  4. 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。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐