Transformer架构解析:为何必须使用多头注意力机制而非单头?

1次阅读
没有评论

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

image.webp

背景痛点

在传统的单头注意力机制中,模型只能学习到一种固定的注意力模式。这在实际应用中会遇到几个关键问题:

Transformer 架构解析:为何必须使用多头注意力机制而非单头?

  • 语义覆盖不足:长序列中不同位置的词语可能涉及多种语义关系(如语法结构、指代关系、情感倾向等),单头注意力难以同时捕获这些多样化特征

  • 梯度传播受限:随着序列长度增加,单头注意力的梯度可能因过度平滑(over-smoothing)而消失,特别是在深层网络中表现更明显

  • 表征瓶颈:所有特征必须通过同一个注意力矩阵压缩,导致信息密度过高时产生特征混淆

机制对比

单头注意力计算公式为:

$$Attention(Q,K,V)=softmax(\frac{QK^T}{\sqrt{d_k}})V$$

而多头注意力的核心改进在于:

  1. 将输入投影到 h 个不同的子空间:
    $$head_i=Attention(QW_i^Q,KW_i^K,VW_i^V)$$

  2. 拼接后二次投影:
    $$MultiHead=Concat(head_1,…,head_h)W^O$$

这种设计带来三个关键优势:

  • 子空间 specialization:每个头可以自主学习不同类型的注意力模式(如局部 / 全局、语法 / 语义)

  • 梯度多样性:不同头的梯度通过独立路径传播,缓解消失问题

  • 容量可控:通过调整头数 h 即可扩展模型能力,而不必大幅增加 d_model

PyTorch 实现

import torch
import torch.nn as nn

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, h=8):
        super().__init__()
        assert d_model % h == 0
        self.head_dim = d_model // h
        self.h = h

        # 投影矩阵(注意实际实现中通常合并成单个大矩阵)self.q_proj = nn.Linear(d_model, d_model)
        self.k_proj = nn.Linear(d_model, d_model) 
        self.v_proj = nn.Linear(d_model, d_model)
        self.out_proj = nn.Linear(d_model, d_model)

    def forward(self, q, k, v, mask=None):
        """
        输入形状: [batch, seq_len, d_model]
        输出形状: [batch, seq_len, d_model]
        """
        batch_size = q.size(0)

        # 线性投影 + 分头 [batch, seq_len, h, head_dim]
        q = self.q_proj(q).view(batch_size, -1, self.h, self.head_dim)
        k = self.k_proj(k).view(batch_size, -1, self.h, self.head_dim)
        v = self.v_proj(v).view(batch_size, -1, self.h, self.head_dim)

        # 转置为方便矩阵运算 [batch, h, seq_len, head_dim]
        q, k, v = q.transpose(1,2), k.transpose(1,2), v.transpose(1,2)

        # Scaled Dot-Product [batch, h, seq_len, seq_len]
        scores = torch.matmul(q, k.transpose(-2,-1)) / (self.head_dim ** 0.5)
        if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)
        attn = torch.softmax(scores, dim=-1)

        # 加权求和 + 转置回原形状 [batch, seq_len, h, head_dim]
        output = torch.matmul(attn, v).transpose(1,2)

        # 拼接所有头 [batch, seq_len, d_model]
        output = output.contiguous().view(batch_size, -1, self.h*self.head_dim)
        return self.out_proj(output)

实验验证

在 IWSLT14 德英翻译任务上的实验结果:

头数 (h) BLEU 训练时间 (epoch)
1 28.3 12.5h
4 31.7 13.1h
8 32.4 14.3h
16 32.1 16.8h

关键发现:

  1. 4- 8 头时达到最佳性能 / 耗时平衡
  2. 头数过多会导致边际效益递减
  3. 当 head_dim<64 时出现训练不稳定(需用梯度裁剪缓解)

生产建议

  1. 维度匹配原则:
  2. 确保 d_model 能被 h 整除(否则需要 padding)
  3. 典型配置:d_model=512 时 h =8(head_dim=64)

  4. 硬件优化:

  5. GPU:使用 tensorcore 时将 h 设为 8 的倍数
  6. TPU:避免 h 超过 128(XLA 编译限制)

  7. 数值稳定性:

  8. 当 head_dim<32 时建议使用:python
    torch.nn.functional.scaled_dot_product_attention()
  9. 或者手动添加 LayerNorm

延伸思考

建议读者尝试:

  1. 可视化不同头的注意力模式(常用方法):

    # 获取第 3 层第 5 头的注意力权重
    attn_weights = model.decoder.layers[2].self_attn.attn[4]
    plt.matshow(attn_weights[0].detach().numpy())

  2. 在业务数据上测试:

  3. 短文本分类:尝试 h =2/4
  4. 长文档摘要:h≥8 效果更好

  5. 可解释性实验:

  6. 固定其他参数,仅改变 h 值观察验证集 loss 变化
  7. 对比不同头学到的 attention pattern 差异
正文完
 0
评论(没有评论)