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

1次阅读
没有评论

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

image.webp

核心概念

单头注意力机制

在标准的单头注意力机制中,输入序列通过三个不同的线性变换得到查询 (Q)、键(K) 和值 (V) 矩阵。假设输入序列长度为 $n$,特征维度为 $d_{model}$,则:

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

  • $Q \in \mathbb{R}^{n \times d_k}$
  • $K \in \mathbb{R}^{n \times d_k}$
  • $V \in \mathbb{R}^{n \times d_v}$

其中 $d_k$ 通常等于 $d_v$,且 $d_k = d_{model}$ 在单头情况下。

注意力得分的计算公式为:
$$Attention(Q, K, V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$

多头注意力机制

多头注意力将 $d_{model}$ 维的 Q、K、V 分别投影到 $h$ 个不同的子空间,每个子空间的维度为 $d_k = d_v = d_{model}/h$。最终将各头的输出拼接后再通过线性变换:

$$MultiHead(Q, K, V) = Concat(head_1,…,head_h)W^O$$
$$where\ head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)$$

痛点分析

单头机制的表征瓶颈

  • 在长序列建模任务中,单头注意力的 Perplexity 指标明显劣于多头机制。实验表明,在 WikiText-103 数据集上,单头比 8 头模型的验证集困惑度高出 15-20%

  • 单头注意力倾向于过度聚焦于局部区域,导致梯度消失。在图像描述生成任务中,单头模型的注意力分布熵值比多头低 30-40%

技术方案

并行化架构

多头注意力的计算过程天然适合并行化,各注意力头可以独立计算。下图展示了维度变换流程:

[输入序列] -> [h 个线性投影] -> [h 个并行 attention 计算] -> [拼接] -> [输出投影]

PyTorch 实现

import torch
import torch.nn as nn
import math

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, h):
        super().__init__()
        assert d_model % h == 0, "d_model must be divisible by h"
        self.d_k = d_model // h
        self.h = h

        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, mask=None):
        batch_size = x.size(0)

        # 线性投影
        Q = self.W_q(x).view(batch_size, -1, self.h, self.d_k).transpose(1,2)
        K = self.W_k(x).view(batch_size, -1, self.h, self.d_k).transpose(1,2)
        V = self.W_v(x).view(batch_size, -1, self.h, self.d_k).transpose(1,2)

        # Scaled Dot-Product Attention
        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 = torch.softmax(scores, dim=-1)

        # 拼接多头输出
        output = torch.matmul(attn, V)
        output = output.transpose(1,2).contiguous().view(batch_size, -1, self.h*self.d_k)

        return self.W_o(output)

性能考量

计算复杂度分析

  • FLOPs 计算公式:$4nd_{model}^2 + 8n^2d_{model}$
  • 内存占用与头数 $h$ 的关系:$MEM \propto h \times (\frac{d_{model}}{h})^2 = \frac{d_{model}^2}{h}$

硬件效率

硬件类型 并行效率(8 头 vs 单头)
GPU 6.8x
TPU 7.2x
CPU 3.5x

避坑指南

  1. 维度验证:确保 $d_{model}$ 能被头数 $h$ 整除

  2. 大模型训练:当使用超过 8 头时,建议启用梯度检查点技术

  3. 混合精度训练:对注意力分数做 $\frac{1}{\sqrt{d_k}}$ 缩放后,再执行 softmax

延伸思考

实验建议

在固定计算量 (FLOPs≈1e9) 条件下,对比两种配置:
– 4 头 256 维
– 8 头 128 维

稀疏注意力

研究发现,将稀疏注意力模式 (如局部窗口注意力) 与多头机制结合时,保留 30-50% 的头做全局注意力效果最佳

参考文献

  1. Vaswani et al. “Attention Is All You Need” (NeurIPS 2017)
  2. Beltagy et al. “Longformer: The Long-Document Transformer” (arXiv 2020)
正文完
 0
评论(没有评论)