模型层 开放阅读

多头注意力

Multi-Head Attention, MHA

概念 ID
multi-head-attention-mha
更新时间
2026-05-29
来源数量
待补

多头注意力

3秒看懂

多头注意力是Transformer架构的“并行理解引擎”——让模型从多个不同表示子空间同时关注输入的不同位置,从而捕获词语间的多重关系(如指代、修饰、语义角色),是当代大语言模型的基石算子。

3分钟产业解释

多头注意力可以通俗地理解为“多人协作审阅一份文件”。单个人阅读可能会遗漏某些细节或偏见;而让多个审阅人(多个“头”)分别用不同的关注点(查询向量)独立审阅,再把他们的发现汇总起来,就得到更全面的理解。在技术实现上,它并行计算多组缩放点积注意力,每组拥有各自的查询(Q)、键(K)、值(V)投影矩阵,最后拼接输出并进行线性变换。多头机制使得模型在不同子空间学习不同类型的依赖关系,例如某个头专门捕捉长距离语法结构,另一个头关注局部语义搭配。产业侧,多头注意力直接决定了大模型训练和推理的算力需求,其计算复杂度为 O(n^2\cdot d)(序列长度 n,维度 d),是 GPU/TPU 设计的核心优化靶点。当前大量硬件(如 NVIDIA H100 的 Transformer Engine)与软件框架(FlashAttention 算法)均围绕高效实现多头注意力而演进。

15分钟专家深入

在现代深度学习栈中,多头注意力不只是算法组件,更是系统设计的枢纽。它的计算图范式深刻影响了并行策略、内存层级和数值精度格式。

并行策略映射:张量并行往往沿着注意力头数维度切分,每个设备持有部分头,然后通过 AllReduce 聚合输出;序列并行则将长序列切片并配合环形注意力(Ring Attention)完成跨设备的键值交互。专家模型(MoE)中的 token 转发使用 All-to-All 通信,但注意力机制的集体通信主要是使用 AllReduce 或 ReduceScatter 等集体通信进行聚合(在张量并行中聚合部分头的输出)。这些通信模式的选择与硬件互联拓扑(NVLink 带宽、节点间 InfiniBand)共同决定训练效率。

算法优化进展:标准 MHA 的空间复杂度 O(b\cdot h\cdot n^2) 促使众多精确近似或 IO 优化出现。FlashAttention 通过分块(tiling)和重计算规避巨型注意力矩阵的 HBM 显存读写,已成为事实标准。多查询注意力(MQA)和分组查询注意力(GQA)通过共享键值投影减少 KV 缓存,大幅降低推理内存占用,已被 LLaMA 2/3 等模型采用。这些变形保留了多头机制的大部分表达能力,同时显著改善服务延迟和吞吐。

数值格式与量化:注意力计算对数值精度敏感。FP8 混合精度训练通常在 Softmax 部分维持高精度(如 FP32),否则梯度会不稳。推理侧,KV 缓存常用 INT8 或 FP8 量化,需要特殊的校准方法(如 SmoothQuant)抑制异常值。2:4 结构化稀疏可在保持硬件加速的同时剪枝一些注意力头,兼顾效能。

关键参数的量级关系:对于千亿参数模型,典型头数在 32–128,每头维度 128,总隐藏维度 d_{model} = h \cdot d_k 通常为 4096 至 16384。序列长度从 2048 扩展到 32768 甚至百万 token(通过位置编码外推),使得注意力矩阵呈平方增长,成为长上下文的关键瓶颈。

技术原理

缩放点积注意力(单头)

给定输入序列表示 X \in \mathbb{R}^{n \times d_{model}},通过权重矩阵 W^Q, W^K, W^V 线性映射得到查询 Q、键 K、值 V,维度 d_k = d_{model}/h

Q = XW^Q, \quad K = XW^K, \quad V = XW^V

注意力权重计算:

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

其中 \sqrt{d_k} 作为缩放因子,防止点积过大导致 Softmax 梯度消失。

多头拼接

h 个头并行执行上述注意力,拼接结果再进行一次线性变换:

\text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1, \dots, \text{head}_h) W^O
\text{where head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)

投影矩阵形状:W_i^Q \in \mathbb{R}^{d_{model} \times d_k}W_i^K \in \mathbb{R}^{d_{model} \times d_k}W_i^V \in \mathbb{R}^{d_{model} \times d_v},通常 d_k=d_v=d_{model}/hW^O \in \mathbb{R}^{hd_v \times d_{model}}

复杂度与内存

  • 计算复杂度:O(n^2 \cdot d_{model})(包括 QK 乘法和加权 V)
  • 内存占用:存储注意力矩阵 n \times n 每个头需要 O(n^2),若缓存中间激活则成为显存瓶颈。标准实现中显存需求 O(b \cdot h \cdot n^2)b 为 batch size。
               输入序列 X: [n, d_model]
                      |
        +-------------+-------------+
        |             |             |
    Linear(W_Q)   Linear(W_K)   Linear(W_V)
        |             |             |
    Q: [n,dk]    K: [n,dk]    V: [n,dv]
        |             |             |
        +------> MatMul(Q,K^T) / sqrt(dk)
                      |
              Softmax(axis=-1)
                      |
                  MatMul(Score, V)
                      |
                 head_output: [n, dv]
        (重复h次,拼接) -> [n, h*dv]
                      |
                 Linear(W_O) -> [n, d_model]

掩码注意力

解码器中的因果掩码(上三角矩阵置 -\infty)确保位置 i 只能看到 i 及之前的信息,维持自回归生成。

技术演进史

  • 2014,基础注意力:Bahdanau 等人在 RNN 编解码器中引入加法注意力,让解码器动态选择源句相关部分。计算代价高,序列化明显。
  • 2015,乘法注意力:Luong 提出点积注意力,简化计算,但仍与 RNN 耦合。
  • 2017,Transformer 与多头注意力:Vaswani 等发布《Attention Is All You Need》,完全抛弃循环和卷积,提出多头缩放点积注意力作为唯一特征交互机制,定义标准 MHA 范式,开启大模型时代。
  • 2019–2020,效率优化:Reformer 使用局部敏感哈希(LSH)注意力降低复杂度;Longformer、BigBird 结合滑动窗口、全局和随机稀疏注意力;Linformer 将注意力矩阵低秩分解,理论复杂度降至 O(n)
  • 2021,多查询注意力(MQA):将所有头共享一副键值投影,大幅减少推理 KV 缓存,PaLM、LaMDA 等模型采用。
  • 2022–2023,分组查询注意力(GQA):介于 MHA 和 MQA 之间,将头分为若干组共享 KV 投影,平衡效率与效果,LLaMA 2、Mistral 等主流模型采用。
  • 2022 至今,IO 感知精确注意力:FlashAttention(Dao et al.)利用 GPU 共享内存和在线 Softmax 减少 HBM 访问,经历 V1/V2/V3,支持 FP8 和动态序列,已成为 Transformer 构建标准背板。
  • 2024+,长上下文与系统融合:Ring Attention 将序列分块并环通信,实现跨设备百万 token 上下文;Striped Attention 异步重叠通信与计算;MLA(Multi-head Latent Attention)在 DeepSeek 等模型中应用,通过低秩键值联合压缩进一步降低推理代价。

技术路线对比(量化表)

以下对比基于公开论文与业界实践(符号:← 越大越好 / 越小越好,定性标记)。

变种头数/分组每头维度KV 缓存大小(相对 MHA)推理显存占用(相对)训练吞吐(相对)长文本质量损失典型模型
MHAh 头,独立 KVd/h低(基准)BERT,GPT-2,ViT
MQAh 头,全部头共享 KV同上1/h极低1.2–1.5×中等PaLM,LaMDA
GQAg 组,组内共享 KV同上g/h1.1–1.3×低–中LLaMA 2/3,Mistral
FlashAttention同 MHA(精确等价)~相同相同(但显存带宽占用低)2–3×(长序列)零(等价)几乎所有现代大模型

注:具体数字因序列长度、硬件而异,[根据公开基准测试估算],未精确披露完整型号的 head-to-head 测量数据。

上下游

  • 上游
    • 基础数学库(cuBLAS、rocBLAS)提供高效矩阵乘法;
    • 底层硬件(GPU Tensor Core、TPU MXU)支持半精度矩阵乘累加;
    • 编译器与框架(Triton、TVM、XLA)生成优化的融合核;
    • 并行策略库(NCCL、NVSHMEM)为头并行提供通信原语。
  • 下游
    • 自然语言处理:语言模型(GPT、LLaMA、Gemini)、机器翻译、文本摘要;
    • 计算机视觉:ViT、Swin Transformer 等将图像分块作为序列输入;
    • 多模态:CLIP、Flamingo、ImageBind 使用交叉注意力对齐图文特征;
    • 语音/生物序列:Whisper、AlphaFold 利用注意力建模时序和残基关系。

关键指标

  • 效率指标
    • TFLOPS(每秒浮点运算次数)在注意力实现上的利用率;
    • 显存带宽利用率(HBM 读写占比);
    • 推理阶段首个 token 延迟与每个 token 生成延迟;
    • KV 缓存容量(字节):约 2\cdot L\cdot h\cdot d_k\cdot N\cdot \text{precision}L 层数,N 序列长度。
  • 能力指标
    • 有效上下文长度:注意力机制能维持准确性的最大 token 距离;
    • 检索精度(如 Needle-in-a-Haystack 测试分数);
    • 不同头功能的可解释性(头重要性、激活稀疏度)。

供需与市场数据

由于未获得最新的专门针对多头注意力市场规模的独立报告(该组件通常被归入更大范围的 AI 训练/推理芯片与软件市场),根据行业趋势定性:随着大型语言模型参数量突破万亿、上下文窗口向 1M token 延伸,高效的多头注意力实现成为 AI 加速器的核心竞争力。据 NVIDIA FY2024 财报披露,数据中心 GPU 营收约 475 亿美元(2024 财年),其中很大比例由 Transformer 及注意力计算驱动。训练成本方面,据公开技术报告估算,在大型 MoE 模型中,注意力计算的浮点运算占比通常显著低于 50%,大部分浮点运算消耗在 FFN/专家层。推理侧,KV 缓存优化能直接降低 30% 以上的 GPU 实例成本,因此 MQA/GQA 等变种已快速普及。开源社区对 FlashAttention 的依赖度极高,几乎所有的 PyTorch / JAX 训练脚本默认调用优化的注意力后端。

代表公司与资本映射

  • 谷歌(Google):多头注意力原提出者,Transformer 架构至今仍是其 Gemini 系列模型的基础;TPU 芯片专门针对注意力计算优化了脉动阵列。
  • 英伟达(NVIDIA):通过 TensorRT-LLM、cuDNN 和 Transformer Engine 提供高性能多头注意力实现,硬件 Blackwell 架构引入专门的注意力引擎;其数据中心 GPU 生态受益于模型规模扩大带来的注意力算力需求激增。
  • OpenAI / Anthropic / Meta 等模型厂商:在内部训练框架中深度定制注意力核(如提升长上下文效率的算法),是 MQA/GQA 等变种的采纳者和推动者。
  • AI 芯片初创(Cerebras、Groq、SambaNova):各自通过数据流架构、大规模 SRAM 或特定算子重构来降低注意力延迟,资本市场高度关注其在长上下文推理中的突破。
  • 风险投资趋势:投资者倾向于支持那些能显著降低 Transformer 注意力计算成本的基础设施公司,包括 FlashAttention 商业化变体(如数据库索引型注意力)、内存池化方案以及存内计算架构。

投资逻辑

  1. 算力供需剪刀差:序列长度每翻倍,裸注意力计算量增长 4 倍;而硬件算力增长放缓(摩尔定律瓶颈),导致高效注意力实现成为刚性价值点,相关软件(FlashAttention 类库)和硬件(支持更优稀疏/量化的芯片)受益。
  2. 推理成本决定模型商业闭环:PaaS / SaaS 部署大模型时,注意力推理占用 30–70% 的 GPU 时间,因此 GQA、KV 缓存量化等技术的成熟度直接影响模型盈利性,押注相关技术栈的公司具备降本增效的长期需求。
  3. 专利与标准风险:Transformer 及 MHA 相关基础专利(如谷歌的注意力机制专利)申请于2018年,仍在20年有效期内,远未过期,但特定优化实现(如 FlashAttention 的 tiling 策略)可能受版权/专利影响,需要关注开源生态的健康度。
  4. 边缘智能:小模型(≤7B)在端侧部署时,内存约束极为苛刻,推动分组注意力、INT4 KV 缓存等方案,利好低功耗存算一体芯片和极致压缩工具链。

常见误读纠偏

误读1:“多头注意力就是让模型同时理解不同含义,头数越多效果越好” 实际上头数受限于每头维度 d_k 不能过小,否则投影矩阵的容量不足以形成有效的注意力子空间。实验表明,适中的头数(如 GPT-3 使用 96 头)已在各项任务上饱和,继续增加头数会带来更多计算开销且收益递减。同时,许多头在训练后表现为冗余或可剪枝,并非所有头都承载独立语种功能,其作用往往是高度混合的。

误读2:“FlashAttention 近似了注意力计算,因此会损失精度” FlashAttention 及其后续版本实现的是数学上完全等价的标准缩放点积注意力,通过分块和在线 Softmax 在不改变数值结果的前提下优化 IO 模式。其输出的梯度与原始实现一致(相对于数值误差环境),因此严格意义上不存在精度损失,不应与其他近似稀疏/低秩注意力混为一谈。

学习路径

  1. 基础入门:阅读《Attention Is All You Need》论文,重点理解 Section 3.2 多头注意力的公式和图形。结合 PyTorch 官方 nn.MultiheadAttention 源码(或最小实现)手写一个 MHA 模块。
  2. 核心算法:学习《FlashAttention: Fast and Memory-Efficient Exact Attention》,理解 Triton 或 CUDA 如何分块计算 Softmax 并减少 HBM 读/写。尝试分析 HuggingFace Transformers 中 LlamaAttention 类的分组查询实现。
  3. 系统与并行:阅读 DeepSpeed Ulysses 序列并行论文,了解张量并行中如何沿着头维度切分注意力;阅读 vLLM 的 PagedAttention 机制,理解 KV 缓存管理如何影响多请求服务。
  4. 产经视角:跟踪 SemiAnalysis、Next Platform 等分析文章,了解不同 AI 芯片(如 Groq LPU,NVIDIA H200)在处理注意力时的带宽和延迟瓶颈。调研开源模型结构报告(如 LLaMA 论文、Mixture-of-Experts 技术报告)中的注意力配置选择及其工程原因。

一句话总结

多头注意力以“多视角并行匹配”取代了序列循环,是 Transformer 大模型的感知中枢,其计算效率与内存方案直接定义了大模型的商用成本天花板。

延伸阅读与来源

  • 原始论文:Vaswani, A. et al. “Attention Is All You Need.” NeurIPS 2017.
  • FlashAttention 系列:Dao, T. et al. “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.” NeurIPS 2022; “FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning.” 2023; “FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision.” 2024.
  • 分组查询注意力:Ainslie, J. et al. “GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints.” EMNLP 2023.
  • 系统实现:NVIDIA Megatron-LM 张量并行论文(Shoeybi et al., 2019); DeepSpeed Ulysses 序列并行; vLLM PagedAttention (Kwon et al., SOSP 2023).
  • 行业分析:受限于检索状态,具体供需数字未直接获取,可关注 Omdia、IDC 关于 AI 服务器/芯片的数据;模型效率分析见 SemiAnalysis 的公开文章。
source: 公开披露与公开资料整理 本页仅用于产业链学习、信息检索和研究辅助;不构成投资建议,不预测涨跌,不提供买卖、仓位或目标价建议。
完整概念页 复盘 13 节结构 公司投研页 沿产业链找到受益公司 投资课 把概念转成可跟踪模型