从零理解Transformer核心:自注意力机制与多头注意力机制的实现原理与实战

1次阅读
没有评论

共计 2772 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

传统 RNN 的困境与注意力机制的诞生

在自然语言处理领域,循环神经网络 (RNN) 曾是处理序列数据的标准方案。但实践中我们发现三个主要问题:

从零理解 Transformer 核心:自注意力机制与多头注意力机制的实现原理与实战

  • 长程依赖丢失:随着序列长度增加,RNN 难以有效传递早期信息(梯度消失 / 爆炸问题)
  • 顺序计算瓶颈:必须按时间步逐步计算,无法利用现代 GPU 的并行能力
  • 固定编码局限:每个时间步的隐藏状态被迫包含所有历史信息,缺乏重点聚焦

自注意力机制 (Self-Attention) 原理详解

核心计算流程

  1. 输入表示
    对于输入序列 $X \in \mathbb{R}^{n \times d_{model}}$(n 为序列长度,$d_{model}$ 为特征维度),通过三个可学习矩阵投影得到:
    $$
    Q = XW^Q, \quad K = XW^K, \quad V = XW^V
    $$
    ($W^Q, W^K \in \mathbb{R}^{d_{model} \times d_k}$, $W^V \in \mathbb{R}^{d_{model} \times d_v}$)

  2. 注意力权重计算
    $$
    \text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V
    $$

  3. 除以 $\sqrt{d_k}$ 防止点积数值过大导致 softmax 梯度消失
  4. softmax 沿每一行计算,保证权重和为 1

  5. 代码实现

    import torch
    import torch.nn as nn
    import torch.nn.functional as F
    
    class SelfAttention(nn.Module):
        def __init__(self, d_model, d_k, d_v):
            super().__init__()
            self.W_q = nn.Linear(d_model, d_k)  # Query 投影
            self.W_k = nn.Linear(d_model, d_k)  # Key 投影
            self.W_v = nn.Linear(d_model, d_v)  # Value 投影
            self.scale = d_k ** -0.5
    
        def forward(self, x):
            """
            输入: [batch_size, seq_len, d_model]
            输出: [batch_size, seq_len, d_v]
            """
            Q = self.W_q(x)  # [B, L, d_k]
            K = self.W_k(x)  # [B, L, d_k]
            V = self.W_v(x)  # [B, L, d_v]
    
            attn = torch.matmul(Q, K.transpose(-1, -2)) * self.scale
            attn = F.softmax(attn, dim=-1)
            output = torch.matmul(attn, V)
            return output

多头注意力机制 (Multi-Head Attention) 进阶

设计动机

  • 单一注意力头的局限:只能学习一种注意力模式
  • 并行化优势:多个头可同时捕捉不同子空间的语义关系

实现关键

  1. 头部拆分与合并
  2. 将 Q /K/ V 拆分为 h 份(h 为头数),每份维度 $d_k=d_v=d_{model}/h$
  3. 计算 h 个独立的注意力头后拼接结果

  4. PyTorch 实现

    class MultiHeadAttention(nn.Module):
        def __init__(self, d_model, num_heads):
            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)
    
        def forward(self, x):
            batch_size = x.size(0)
    
            # 投影后拆分多头 [B, L, d_model] -> [B, L, h, d_k]
            Q = self.W_q(x).view(batch_size, -1, self.num_heads, self.d_k)
            K = self.W_k(x).view(batch_size, -1, self.num_heads, self.d_k)
            V = self.W_v(x).view(batch_size, -1, self.num_heads, self.d_k)
    
            # 转置为 [B, h, L, d_k]
            Q = Q.transpose(1, 2)
            K = K.transpose(1, 2)
            V = V.transpose(1, 2)
    
            # 计算缩放点积注意力
            scores = torch.matmul(Q, K.transpose(-1, -2)) / (self.d_k ** 0.5)
            attn = F.softmax(scores, dim=-1)
            context = torch.matmul(attn, V)
    
            # 合并多头 [B, h, L, d_k] -> [B, L, d_model]
            context = context.transpose(1, 2).contiguous()
            context = context.view(batch_size, -1, self.num_heads * self.d_k)
    
            return self.W_o(context)

性能对比与实验观察

指标 单头注意力 多头注意力 (h=8)
计算复杂度 O(n²d) O(n²d)
参数量 3dd_k 3dd
并行度
语义捕获能力 单一模式 多样化模式

实际任务中(如机器翻译),多头注意力通常能带来 1.5-2.5 BLEU 值提升。

五大避坑指南

  1. 维度不对齐错误
  2. 现象:RuntimeError: mat1 and mat2 shapes cannot be multiplied
  3. 检查:确保 $d_{model}$ 能被头数整除,投影后张量形状匹配

  4. softmax 数值溢出

  5. 现象:注意力权重出现 NaN
  6. 解决:必须进行缩放(除以 $\sqrt{d_k}$),对特别长的序列可分段计算

  7. 梯度消失问题

  8. 现象:模型难以学习远程依赖
  9. 对策:配合残差连接和 LayerNorm 使用

延伸思考方向

  1. 头数超参数选择
  2. 实验发现不同任务最优头数不同(翻译常用 8 头,分类可能 4 头足够)
  3. 可通过注意力头可视化分析各头学习到的模式

  4. 计算效率优化

  5. 稀疏注意力、局部注意力等变体如何权衡效果与速度
  6. FlashAttention 等优化技术原理探究

总结启示

通过实现完整的自注意力和多头注意力模块,我们深入理解了 Transformer 的核心设计思想。关键收获包括:
– 注意力机制通过动态权重实现序列元素的直接交互
– 多头设计类似 CNN 的多通道,能并行学习多样特征
– 实际使用时需注意数值稳定性和计算效率的平衡

建议读者在完成基础实现后,进一步尝试将其应用到具体 NLP 任务中,观察不同超参数配置下的性能变化,这将大大加深对机制的理解。

正文完
 0
评论(没有评论)