从零实现Transformer核心组件:位置编码与Mask机制的PyTorch实战指南

当你第一次翻开Transformer论文时,那些复杂的数学公式是否让你望而却步?作为自然语言处理领域的革命性架构,Transformer的核心思想其实可以通过代码直观理解。本文将带你用PyTorch从零实现两个关键组件——位置编码和Mask机制,通过编写可运行的代码来反向理解其设计原理。

1. 位置编码:让模型记住词序的魔法

传统RNN天然具有处理序列的能力,而Transformer作为无时序的架构,需要额外机制来编码位置信息。这就是位置编码(Positional Encoding)的用武之地。

1.1 为什么需要位置编码?

想象你在阅读这句话:"猫追老鼠"和"老鼠追猫"。词序不同,语义完全相反。Transformer的自注意力机制本身无法感知词序,因此需要显式地注入位置信息。

位置编码的设计需要满足几个关键特性:

  • 唯一性 :每个位置有唯一编码
  • 相对位置关系 :能够表示位置间的相对距离
  • 泛化性 :能处理比训练时更长的序列

1.2 正弦余弦编码的数学实现

原始论文使用交替的正弦和余弦函数来生成位置编码。对于位置pos和维度i,编码公式为:

PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))

这种设计有几个精妙之处:

  • 正弦函数的 周期性 让模型能学习到相对位置关系
  • 不同维度使用不同的 波长 ,形成多层次的位置感知
  • 数值范围被控制在[-1,1]之间,与词嵌入尺度匹配

1.3 PyTorch完整实现

让我们用PyTorch实现这一机制:

import torch
import math

class PositionalEncoding(torch.nn.Module):
    def __init__(self, d_model: int, max_len: int = 5000):
        super().__init__()
        
        # 创建位置编码矩阵
        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
        
        self.register_buffer('pe', pe)  # 不参与训练

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Args:
            x: Tensor, shape [batch_size, seq_len, embedding_dim]
        """
        x = x + self.pe[:x.size(1)]  # 只取前seq_len个位置
        return x

这段代码的几个关键点:

  1. 使用矩阵运算而非循环,大幅提升效率
  2. register_buffer 确保位置编码不参与训练
  3. 直接加到词嵌入上,实现位置信息注入

提示:实际应用中,位置编码通常与词嵌入相加而非拼接,这能减少模型参数同时保持信息融合

1.4 位置编码的可视化分析

让我们观察生成的位置编码矩阵:

import matplotlib.pyplot as plt

d_model = 512
max_len = 100
pe = PositionalEncoding(d_model, max_len).pe

plt.figure(figsize=(12, 6))
plt.imshow(pe[:100, :100].T, cmap='viridis')
plt.xlabel('Position')
plt.ylabel('Dimension')
plt.colorbar()
plt.show()

你会看到:

  • 低维度(图底部)变化剧烈,捕获细粒度位置信息
  • 高维度(图顶部)变化平缓,捕获宏观位置关系
  • 相邻位置编码相似但有差异,形成平滑过渡

2. Mask机制:控制信息流的阀门

Transformer中有两种关键Mask:填充Mask(pad_mask)和序列Mask(tril_mask)。它们分别解决不同问题。

2.1 填充Mask:处理变长序列

在实际应用中,我们通常将多个句子打包成批次(batch)处理。由于句子长度不一,需要填充(pad)到相同长度。填充Mask确保模型忽略这些无意义的填充位置。

实现原理

  1. 标记所有填充位置为True
  2. 将这些位置对应的注意力权重设为负无穷
  3. 经过softmax后,这些位置的权重变为0
def create_pad_mask(seq: torch.Tensor, pad_idx: int) -> torch.Tensor:
    """
    创建填充mask
    Args:
        seq: 输入序列,形状 [batch_size, seq_len]
        pad_idx: 填充token的索引
    Returns:
        mask: 形状 [batch_size, 1, 1, seq_len]
    """
    mask = (seq == pad_idx).unsqueeze(1).unsqueeze(2)  # [batch_size, 1, 1, seq_len]
    return mask

2.2 序列Mask:防止信息泄露

在解码器中,预测下一个词时不应看到"未来"的信息。序列Mask(又称因果Mask)通过上三角矩阵实现这一点。

def create_seq_mask(seq_len: int) -> torch.Tensor:
    """
    创建序列mask(上三角矩阵)
    Args:
        seq_len: 序列长度
    Returns:
        mask: 形状 [1, seq_len, seq_len]
    """
    mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()
    return mask.unsqueeze(0)  # [1, seq_len, seq_len]

注意:在训练解码器时,需要同时使用填充Mask和序列Mask,通过逻辑或操作合并它们

2.3 组合Mask的实现

实际应用中,我们经常需要组合多种Mask:

def combine_masks(pad_mask: torch.Tensor, seq_mask: torch.Tensor) -> torch.Tensor:
    """
    组合填充mask和序列mask
    Args:
        pad_mask: 形状 [batch_size, 1, 1, seq_len]
        seq_mask: 形状 [1, seq_len, seq_len]
    Returns:
        组合后的mask: 形状 [batch_size, 1, seq_len, seq_len]
    """
    if pad_mask is None and seq_mask is None:
        return None
    
    combined_mask = pad_mask | seq_mask if pad_mask is not None and seq_mask is not None else \
                   pad_mask if pad_mask is not None else seq_mask
    
    return combined_mask

3. 注意力机制中的Mask应用

理解了Mask的原理后,我们来看如何在注意力机制中应用它们。

3.1 缩放点积注意力实现

def scaled_dot_product_attention(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    mask: torch.Tensor = None
) -> torch.Tensor:
    """
    缩放点积注意力计算
    Args:
        q: query, 形状 [batch_size, n_heads, seq_len, d_k]
        k: key, 形状 [batch_size, n_heads, seq_len, d_k]
        v: value, 形状 [batch_size, n_heads, seq_len, d_v]
        mask: 可选的mask, 形状 [batch_size, 1, seq_len, seq_len]
    Returns:
        注意力输出: 形状 [batch_size, n_heads, seq_len, d_v]
    """
    d_k = q.size(-1)
    attn_scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k)
    
    if mask is not None:
        attn_scores = attn_scores.masked_fill(mask == 1, float('-inf'))
    
    attn_weights = torch.softmax(attn_scores, dim=-1)
    output = torch.matmul(attn_weights, v)
    
    return output

3.2 多头注意力中的Mask处理

在完整的Transformer实现中,Mask需要适配多头注意力:

class MultiHeadAttention(torch.nn.Module):
    def __init__(self, d_model: int, n_heads: int):
        super().__init__()
        assert d_model % n_heads == 0, "d_model必须能被n_head整除"
        
        self.d_model = d_model
        self.n_heads = n_heads
        self.d_k = d_model // n_heads
        
        self.w_q = torch.nn.Linear(d_model, d_model)
        self.w_k = torch.nn.Linear(d_model, d_model)
        self.w_v = torch.nn.Linear(d_model, d_model)
        self.w_o = torch.nn.Linear(d_model, d_model)
        
    def forward(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        mask: torch.Tensor = None
    ) -> torch.Tensor:
        batch_size = q.size(0)
        
        # 线性投影
        q = self.w_q(q).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        k = self.w_k(k).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        v = self.w_v(v).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        
        # 应用mask(如果需要)
        if mask is not None:
            mask = mask.unsqueeze(1)  # 适配多头 [batch_size, 1, 1, seq_len] -> [batch_size, 1, 1, seq_len, seq_len]
        
        # 计算注意力
        attn_output = scaled_dot_product_attention(q, k, v, mask)
        
        # 合并多头
        attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
        
        return self.w_o(attn_output)

4. 实战:构建简易Transformer层

现在我们将位置编码和Mask机制整合到一个简易的Transformer层中。

4.1 完整层实现

class TransformerLayer(torch.nn.Module):
    def __init__(self, d_model: int, n_heads: int, ff_dim: int, dropout: float = 0.1):
        super().__init__()
        self.self_attn = MultiHeadAttention(d_model, n_heads)
        self.ffn = torch.nn.Sequential(
            torch.nn.Linear(d_model, ff_dim),
            torch.nn.ReLU(),
            torch.nn.Linear(ff_dim, d_model)
        )
        self.norm1 = torch.nn.LayerNorm(d_model)
        self.norm2 = torch.nn.LayerNorm(d_model)
        self.dropout = torch.nn.Dropout(dropout)
        
    def forward(self, x: torch.Tensor, mask: torch.Tensor = None) -> torch.Tensor:
        # 自注意力
        attn_output = self.self_attn(x, x, x, mask)
        x = x + self.dropout(attn_output)
        x = self.norm1(x)
        
        # 前馈网络
        ffn_output = self.ffn(x)
        x = x + self.dropout(ffn_output)
        x = self.norm2(x)
        
        return x

4.2 使用示例

# 参数设置
d_model = 512
n_heads = 8
ff_dim = 2048
seq_len = 100
batch_size = 32
vocab_size = 10000

# 初始化组件
embedding = torch.nn.Embedding(vocab_size, d_model)
pos_encoder = PositionalEncoding(d_model)
transformer_layer = TransformerLayer(d_model, n_heads, ff_dim)

# 模拟输入
input_seq = torch.randint(0, vocab_size, (batch_size, seq_len))
pad_mask = create_pad_mask(input_seq, pad_idx=0)  # 假设0是pad索引
seq_mask = create_seq_mask(seq_len)

# 前向传播
x = embedding(input_seq)
x = pos_encoder(x)
output = transformer_layer(x, mask=combine_masks(pad_mask, seq_mask))

4.3 调试技巧

在实现Transformer组件时,有几个调试技巧很有用:

  1. 形状检查 :在每个关键步骤打印张量形状

    print(f"Shape after embedding: {x.shape}")
    
  2. Mask验证 :可视化Mask确保其正确性

    plt.imshow(pad_mask[0, 0, 0].cpu().numpy())
    plt.title('Pad Mask')
    plt.show()
    
  3. 梯度检查 :使用PyTorch的autograd.gradcheck验证反向传播

    from torch.autograd import gradcheck
    test_input = torch.randn(2, 10, d_model, dtype=torch.double, requires_grad=True)
    test = gradcheck(TransformerLayer(d_model, n_heads, ff_dim), test_input)
    print("Gradient check passed:", test)
    

通过本文的代码实现,你应该对Transformer的位置编码和Mask机制有了更直观的理解。这些核心组件不仅是Transformer成功的关键,也是理解现代NLP模型的基础。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐