用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"(狗跑得快)时,模型会:

  1. 首先关注主语"der Hund"
  2. 然后聚焦动词"läuft"
  3. 最后处理副词"schnell"

这种注意力模式与人类翻译过程惊人地相似。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐