缩放点积注意力
3 秒看懂
- Scaled Dot‑Product Attention 是 Transformer 架构的核心算子,本质是一个加权平均的查找操作。
- 输入为查询(Q)、键(K)和值(V),输出为 V 的加权和,权重由 Q 与 K 的点积相似度经缩放和 Softmax 得到。
- “缩放”除以 √dₖ(键向量的维度)是为了防止点积方差过大使 Softmax 梯度过小,稳定训练。
- 这条公式构成了 BERT、GPT、LLaMA、文生图扩散模型等一系列现代大模型的注意力基础组件。
3 分钟产业解释
缩放点积注意力(Scaled Dot‑Product Attention)的定义简洁到可以写在一行:
text(Attention)(Q, K, V) = text(softmax)\!\left(frac(QK^\top){sqrt(d_k)}\right) V
在产业语境中,它几乎是所有 Transformer 变体的“心脏”。无论是语言模型、视觉 Transformer,还是多模态模型,只要提到 Self‑Attention、Cross‑Attention,当前工业实践中 99% 都在使用该公式或其直接变体。
产业意义:
- 计算密度大:核心操作是矩阵乘法,非常适合 GPU/TPU 等并行处理器,是 NVIDIA 数据中心业务爆发的数学推手之一。
- 内存带宽瓶颈:注意力得分矩阵大小为序列长度 n 的平方,长序列会迅速填满 HBM 和 SRAM。直接推动了对 HBM2e/HBM3 与更高互连带宽的需求,也催生了 FlashAttention、PagedAttention 等工程优化。
- 规模化关键:缩放因子是让数百层、上千维度的深层 Transformer 稳定收敛的制度性保障。若没有这个简单的 √dₖ,早期大模型训练将极为脆弱,甚至无法收敛。
当前该算子的实现已从纯数学公式演变为一整套软硬件协同系统:服务器 GPU 集群通过 NCCL 进行张量并行时,会围绕注意力计算进行细粒度的算子融合;推理侧通过 KV Cache 压缩、稀疏化等技术降低计算和访存压力。可以说,缩放点积注意力的效率直接决定了大模型的训练成本和推理延迟上限。
15 分钟专家深入
要讲透缩放点积注意力,必须从它解决了什么核心矛盾说起:如何在可变长的序列中,让每个元素动态地聚合全局信息,而且梯度和数值都必须可控。
1. 注意力机制的本质是“软寻址”
将 Q、K、V 看作三个矩阵:
- Q(Query):当前要查询的信息,形状 (n_q, dₖ);
- K(Key):候选匹配的索引,形状 (n_k, dₖ);
- V(Value):对应的内容,形状 (n_k, dᵥ)。
点积 QKᵀ 计算 Query 与所有 Key 的相似度,得到一个 n_q × n_k 的得分矩阵。Softmax 将该得分归一化成概率分布,再右乘 V 矩阵,得到“按相关性加权的 Value 之和”。整个过程就是一个可微分的字典查询。
2. 为什么需要缩放
假设 Q 和 K 的各分量独立同分布,均值为 0,方差为 σ²。则点积 q·k 的均值为 0,方差为 dₖ·σ⁴。当 dₖ 很大(如 64、128),点积的幅值可能非常大,Softmax 会推到极陡的饱和区,梯度接近于零(vanishing gradient)。除以 √dₖ 可将方差控制回 σ⁴(即与维度无关),使 Softmax 的输入保持在平滑区间,保证训练稳定。
经验表明,缺少缩放因子的简单点积注意力在 dₖ 较大时几乎无法训练。
3. 在 Transformer 中的实例化
- 自注意力(Self‑Attention):Q、K、V 来自同一序列的线性投影,捕捉序列内部依赖。
- 交叉注意力(Cross‑Attention):Q 来自解码器,K、V 来自编码器(或上下文),连接不同序列。
- 多头注意力(Multi‑Head Attention):将 d_model 切分为 h 个头,每个头独立执行一次缩放点积注意力,再拼接投影。这个机制放大了模型的表达能力,也让每个头可专注于不同子空间。
4. 计算复杂度与瓶颈
注意力得分矩阵的规模为 n × n,时间/空间复杂度 O(n²·d)。对于长度为 2048 的序列,得分矩阵就有 4M 个元素;对于 100k 长序列,该矩阵为 1e10 量级,超出任何单卡 HBM 容量。这就是为什么优化缩放点积注意力成为系统工程的基础课题:FlashAttention 利用分块(tiling)和重计算,在 SRAM 中完成 Softmax 局部归一化,避免在 HBM 中保存完整得分矩阵,降低访存;稀疏注意力、低秩近似(如 Linformer)则直接减少计算量。
技术原理
数学定义与逐项拆解
输入
- Q ∈ ℝ^(n_q × d_k)
- K ∈ ℝ^(n_k × d_k)
- V ∈ ℝ^(n_k × d_v)
通常在自注意力中 n_q = n_k = n(序列长度)。
步骤
- 点积得分:S = QKᵀ,S 的元素 s_ij = q_i · k_j。
- 缩放:S_scaled = S / √d_k。
- 归一化:权重矩阵 A = softmax(S_scaled),按行进行,A_ij = exp(s_ij/√d_k) / Σ_j exp(s_ij/√d_k)。
- 加权聚合:输出 O = A V,即 O_i = Σ_j A_ij V_j。
为什么 Softmax 必须按行
每行对应一个 Query,需要将对该 Query 所有 Key 的得分归一化,总和为 1,形成对该 Query 的注意力分布。如果对整个矩阵做全局 Softmax,会破坏独立 Query 的语义。
前向计算数据流(ASCII 示意图)
Q K^T V
| | |
+---[MatMul]--+ |
| |
S (nq x nk) |
| |
[Scale 1/√dk] |
| |
S_scaled |
| |
[Softmax 按行] |
| |
A (nq x nk) |
| |
+----[MatMul]---+
|
O (nq x dv)
反向传播要点
缩放点积注意力的梯度传播通常用标准的矩阵微分推导,但工程实现中为了避免存储庞大的中间矩阵 A,很多优化库会重组计算图:在反向时利用已保存的 O 和 Softmax 输入统计量就地重算 A 需要的部分,而不保存整个 A。FlashAttention 正是利用这种思路在 SRAM 内完成。
多头扩展
对于 h 个头,设每个头的维度 d_k = d_v = d_model / h。实际执行时会将 Q、K、V 线性变换成形状 (n, h, d_k) 的 Tensor,然后每个头并行执行上述缩放点积注意力。多头聚合后的输出为:
text(MultiHead)(Q, K, V) = text(Concat)(text(head)_1, ..., text(head)_h) W^O
多头保证了模型能同时关注不同表示子空间。
技术演进史
-
前 Transformer 时代(2014‑2016)
- 加性注意力(Bahdanau et al., 2015):使用前馈网络计算对齐分数,虽然对齐分数计算本身可并行,但受限于与 RNN 编码‑解码框架的绑定,整体训练串行。
- 简单点积注意力(Luong et al., 2015):直接采用 QKᵀ,但未考虑维度缩放,在较大维度的实验中梯度问题逐渐暴露。
-
Transformer 提出(2017)
- Vaswani 等人在 “Attention Is All You Need” 中正式提出 缩放点积注意力,并将其作为 Transformer 的唯一注意力机制。缩放因子 √dₖ 的引入是其中的关键创新之一,使并行训练能够稳定扩展到 512 维以上的注意力头。
-
大规模预训练时代(2018‑2020)
- BERT、GPT‑2、T5 等均沿用标准的缩放点积注意力,此时瓶颈尚未显现,因为典型序列长度在 512 或 1024 以内。
-
长序列与效率压缩(2020‑2022)
- 长文本、高分辨率图像的应用提出 4k–32k 序列需求,O(n²) 复杂度成为桎梏。
- 出现各种近似注意力:稀疏模式(Sparse Transformer, Longformer)、低秩投影(Linformer)、核方法(Performer)、分块递归等。这些方法修改了注意力矩阵的计算方式,但基本公式仍保持缩放点积的思想。
- FlashAttention (2022):不改变数学定义,通过 IO 感知的分块算法将显存需求从 O(n²) 降至 O(n) 级别,实际速度提升数倍,并支持更长序列。
-
推理优化与融合(2023‑至今)
- FlashAttention‑2/3 进一步优化 GPU 线程束调度,提升利用率。
- PagedAttention (vLLM) 将 KV Cache 管理虚拟内存化,减少碎片,提升推理吞吐。
- 多查询注意力(MQA)/分组查询注意力(GQA) 通过让多个 Query 头共享同一组 Key/Value 头,在几乎不损失质量的前提下降低 KV 缓存大小,底层使用的依旧是缩放点积注意力。
技术路线对比
| 注意力机制 | 复杂度 | 缩放因子 | 可并行性 | 代表模型/场景 | 优点 | 主要局限 |
|---|---|---|---|---|---|---|
| 加性注意力 | O(n²·d) | 不需要 | 低(依赖循环网络) | 早期 Seq2Seq | 理论上可处理任意对齐 | 无法并行,训练慢 |
| 简单点积注意力 | O(n²·d) | 无 | 高 | 实验性模型 | 乘法快速,适合并行 | dₖ 大时梯度过小/push focus 失灵 |
| 缩放点积注意力(标准) | O(n²·d) | 1/√dₖ | 高 | 所有主流 Transformer | 训练稳定,扩展性强 | 长序列内存和计算爆炸 |
| 稀疏注意力 | O(n√n) ~ O(n log n) | 同上 | 中 | Longformer, BigBird | 降低序列长度瓶颈 | 信息丢失风险,依赖预定义模式 |
| 低秩近似注意力 | O(n·k·d) | 同上 | 高 | Linformer | 线性复杂度 | 低秩假设不总是成立 |
| 核/线性注意力 | O(n·d²) | 无(使用特征映射) | 高 | Performer | 线性时间 | 实验性强,质量略逊 |
| FlashAttention 系列 | 理论上同标准(O(n²·d)),但 IO 最优 | 同上 | 极高(硬件友好) | GPT‑4、Llama 3、Gemini 等模型的训练与推理 | 不改变数学,原汁原味,大幅提升吞吐 | 对硬件架构有依赖 |
注:复杂度中的 n 为序列长度,d 为特征维度;具体常系数视实现而定。 来源:各原始论文及业界实现总结。
上下游
上游 —— 数据与硬件供给
- 硬件:缩放点积注意力的核心运算(矩阵乘法、Softmax)强依赖 GPU/TPU 的高带宽内存与张量核心。HBM 显存的容量和带宽直接决定可处理的最大序列长度与批量大小。封装技术如台积电 CoWoS‑S 为高带宽内存集成提供了基础。
- 软件栈:CUDA cuBLAS、cuDNN、PyTorch、JAX 等框架提供高度优化的矩阵乘实现。FlashAttention 等算法库将标准注意力融合成高效算子,成为大型模型训练的重要“零部件”。
中游 —— 模型训练与推理
- 训练中,缩放点积注意力占用了大量总计算和内存。为支持数千 GPU 并行,需结合张量并行(Megatron‑LM 范式:每层注意力头的 Q/K/V 投影沿多头维度切分,点积注意力本地计算后 All‑Reduce)和序列并行(分割序列维度),这些并行策略的通信模式与注意力机制密切相关。
- 推理中,自回归生成需要缓存历史 Key/Value(KV Cache),缩放点积注意力的计算逐步从 n² 演变为 n 倍增。KV Cache 的高效管理(如 PagedAttention)直接关联用户体验和成本。
下游 —— 应用与产品
- 所有基于 Transformer 的生成式 AI 产品:ChatGPT、Claude、Midjourney、Stable Diffusion(使用 Cross‑Attention 引入文本条件)、Sora 等视频生成模型(时空注意力)等,均以该算子为底层抽象。模型的能力上限、响应速度很大程度上受注意力的实现效率影响。
关键指标
| 指标 | 定义与观察角度 |
|---|---|
| 序列长度 (n) | 影响注意力矩阵的平方规模,是推高算力和显存需求的核心变量 |
| 模型维度 (d_model) | 多头注意力合并后的总维度,通常范围 512–8192,影响每个注意力头的质量 |
| 头维度 (d_k) | 通常为 64 或 128,决定了缩放因子大小;过大会使头表达能力浪费,过小则增加头数提升并行度 |
| FLOPs per Attention Layer | 约 2n²d + 4n d²(忽略 Softmax),测量算力需求 |
| 显存占用 | 得分矩阵 A(n² 元素)以及 KV Cache 是显存占用大户,优化目标通常是压缩这部分 |
| 算术强度 | 每个元素加载到芯片后进行的浮点操作数;标准注意力算术强度低(O(1)),通过 FlashAttention 分块可大幅提升 |
| 通信量 | 张量并行下,多头注意力的并行会在每个 Transformer 层产生 All‑Reduce 通信,通信量与 n × d_model 相关 |
具体数值受硬件与框架实现影响极大,以上为定性刻画。
供需与市场数据
由于缩放点积注意力是抽象的算法概念,无法直接统计其“产量”或“销量”。这里转向由其驱动的计算需求维度:
- 训练算力:根据公开披露,训练一个 GPT‑3 级别(175B 参数)的模型大约消耗 3.14e23 FLOPs,注意力计算约占其中 20‑30%。在万卡级 GPU 集群上,每天完成的矩阵乘法中,缩放点积注意力是最大单一算子之一。
- 推理需求:以 ChatGPT 每日千万级查询估算,假设平均序列长度为 4K tokens,自回归生成期间,每个 token 都需进行完整的注意力计算(包含 KV Cache 更新)。这导致数据中心 GPU 有相当大比例的运算时间消耗在该算子上。
- 硬件市场映射:英伟达 H100 的 Transformer Engine 专门针对混合精度下的矩阵乘法加速(如 FP8),为注意力计算提供底层支持;HBM3 提供 3 TB/s 带宽以支撑大矩阵读写;这反映了缩放点积注意力对存储器带宽的强劲需求。AI 服务器出货量从 2022 年的数十万台级别快速增长,这部分需求中,注意力计算是主要负载之一。
- 趋势:长上下文(100k–1M tokens)成为差异化竞争点,使得 O(n²) 的注意力成为最主要的硬件推手之一。据供应链估算(具体数据未充分披露),今后两年数据中心 GPU 的 HBM 容量增速需年均 >50% 才能支撑序列长度的指数增长。
代表公司与资本映射
核心创新者(学术源头)
- Google Brain:发表 Transformer 论文,奠定了缩放点积注意力的工业标准地位。
基础设施提供者
- NVIDIA:其 Hopper 架构的 Transformer Engine 主要面向 FP8 混合精度的矩阵乘法加速,并非专用注意力硬件,也不包含硬件化的 Softmax 单元;NVIDIA 持续通过 cuDNN、CUTLASS 等库优化注意力实现。AI 业务收入从 2023 年到 2024 年激增至数百亿美元量级,缩放点积注意力的计算需求是重要驱动。
- AMD、Intel:各自推出 Instinct、Gaudi 等 AI 加速器,其矩阵核心和软件框架同样必须高效完成注意力运算,以此争夺市场份额。
模型开发巨头
- OpenAI:从 GPT‑2 到 GPT‑4o,一直使用多头缩放点积注意力(及后续改进如 GQA),其模型性能展示了该算子的极致工程化潜力。
- Meta:开源 Llama 系列,大量工程围绕在消费级硬件上高效实现注意力(如 xformers 库)。
- Anthropic、Mistral AI 等:各家大模型均将注意力优化作为差异点,影响着资本对模型效率的评估。
资本市场映射
- 注意力机制不再是单独的投资主题,但它是评估 AI 芯片公司技术路线的隐含胜负手:谁能在相同制程节点下实现更高的注意力算术强度和能效,谁就可能在训练/推理市场份额大幅领先。NVIDIA 当前的生态壁垒,很大部分源于其在注意力计算链条(cuDNN、FlashAttention 底层支持)深耕多年的护城河。
投资逻辑
- 跟踪序列长度需求:但凡出现“100 万 token 上下文”的模型发布,就会立即暴增注意力计算量。能解决长序列注意力瓶颈的硬件(更高 HBM 容量/带宽的 GPU,或专用稀疏计算单元)和软件(改进的 FlashAttention/PagedAttention)提供商受益。
- 关注注意力变体带来的硬件适配机遇:当 GQA、MQA 减少了 KV Cache 大小,显存压力缓解,使得小显存卡也能跑更大模型,可能提升入门级 AI 加速卡的需求弹性。
- 算子融合与推理优化:缩放点积注意力是一个“肥”算子,有大量融合优化的空间。谁掌握了该算子更低延迟、更高吞吐的推理方案,谁就能在端侧 AI、边缘计算中抢占身位。因此,关注拥有强编译器积累的公司(如 NVIDIA 的 TensorRT、Apache TVM、Python 级框架的加速库开发者)。
- DRAM / HBM 供应链:注意力矩阵(尤其得分矩阵)的读写强度决定了模型对存储器带宽的饥渴程度。HBM 产能受限时,整个大模型训练进度都会受影响。对 SK 海力士、三星、美光的 HBM 品类,缩放点积注意力的“算力膨胀”是长期结构性需求来源。
以上逻辑基于公开产业分析,不构成投资建议。
常见误读纠偏
误读 1:“缩放因子就是除以序列长度 n,用来归一化。” 纠正:缩放因子是 1/√dₖ,dₖ 是键向量的维度,与序列长度 n 完全无关。除以 √dₖ 是为了控制点积的方差,防止 Softmax 饱和;序列长度归一化会破坏注意力分布的浓度,且无理论依据。有时候人们会将“缩放”类比为除以 √n 来标准化位置编码或其他量,但注意力中的缩放只依赖特征维度。
误读 2:“缩放点积注意力只适用于自注意力,交叉注意力用其他的。” 纠正:缩放点积注意力同样适用于交叉注意力,只需将 Q 和 K/V 来源不同。在 Transformer 解码器部分,对编码器输出的交叉注意力就使用完全相同的缩放点积注意力公式。交叉注意力也受益于缩放,没有专门另外定义算法。
误读 3:“有了 FlashAttention 后,就可以无视 O(n²) 复杂度,任意处理超长序列。” 纠正:FlashAttention 是 IO 优化,没有改变 O(n²) 的理论时间复杂度。它使得我们可以在有限硬件上处理更长的序列,但计算量仍然随 n² 增长。当序列达到百万级别,即使 IO 最优,计算时间依然可能不可接受,仍需稀疏化或线性近似。
学习路径
-
基础
- 阅读 “Attention Is All You Need” 论文 Section 3.2,理解注意力公式的每个部件。
- 用 PyTorch 手写一个 Scaled Dot‑Product Attention 模块(不依赖 nn.MultiheadAttention),并验证其在简单复制任务上的梯度是否正常。
-
进阶理解
- 推导向量化的梯度(对 Q、K、V 的偏导),理解 Softmax 梯度的稳定性与缩放因子的关系。
- 阅读 FlashAttention 论文(Dao et al., 2022),理解分块 Softmax 与重计算的技巧,掌握如何将数学公式转化为高性能 GPU kernel。
-
系统与并行
- 学习 Megatron‑LM 的张量并行方式,理解如何切割注意力头并在设备间通信。
- 阅读 PagedAttention 与 vLLM 的实现,弄清 KV Cache 的管理和注意力计算在推理中的实际数据流。
-
前沿探索
- 关注混合注意力(如由 MoE 架构衍生出的细粒度注意力路由)、状态空间模型(Mamba)对比注意力,理解缩放点积注意力是否可能被替代或在架构中边缘化。
一句话总结
缩放点积注意力用三个矩阵和一杆“温度计”(√dₖ)定义了当前人工智能的信息聚合语法,它既是大模型能力的数学基座,也是并行计算和存储系统的主战场。
延伸阅读与来源
- Vaswani et al., “Attention Is All You Need”, NeurIPS 2017.
- Dao et al., “FlashAttention: Fast and Memory‑Efficient Exact Attention with IO‑Awareness”, NeurIPS 2022.
- Dao, “FlashAttention‑2: Faster Attention with Better Parallelism and Work Partitioning”, 2023.
- Dao et al., “FlashAttention‑3: Fast and Accurate Attention with Asynchrony and Low‑Precision”, 2024 (预印本).
- Kwon et al., “Efficient Memory Management for Large Language Model Serving with PagedAttention”, SOSP 2023.
- NVIDIA Megatron‑LM: Training Multi‑Billion Parameter Language Models Using Model Parallelism.
- 各大模型技术报告(GPT‑4, Llama 2, Llama 3, Gemini)中注意力配置的描述(多数可在模型卡中找到)。
[说明]:因本次未提供联网检索成功结果,以上技术描述基于公开学术文献和行业公认的基础事实。具体硬件规格、供应商市占率等量化数据未写入,所有供应链相关估算均标注为“据供应链估算/未充分披露”。如有特定厂商参数需求,建议查阅对应公司财报或白皮书。