跳到主要内容

注意力核心 — SDPA 与三个省法(GQA、滑窗、sink)

这一章讲两件事: 缩放点积注意力(SDPA,transformer 的正式名称)五步到底怎么算; 它「每加长一倍、计算翻四倍」的成本怎么被三个省法按住——GQA、滑动窗口、sink。 这是全书的发动机章:第 07 章把它总装成块,第 09 章码进完整模型。

1. 先看现象:为什么长文这么贵

注意力要「每个词看所有词」。这句话的另一面是:文字串翻倍,词对数翻四倍—— 10 万 token 的文档,注意力要算 100 亿个词对。

原书把这个账写成复杂度(规模变大时成本怎么涨的记法):O(n²),其中 n 是这一串的长度。

这种一长串 token,行话叫序列——每长一点,账单平方着涨1。 第 09 章还会算显存账,这里先立住一个问题感:注意力好用,但全量注意力贵到用不起。

原书把这一章讲的机制统称 SDPA(缩放点积注意力):「点积」指打分用向量点积(对应位置相乘再加总), 「缩放」指打完分要除一个数,「注意力」指分数归一化后当关注度用2。 它是 2017 年那篇论文对更早「加性注意力」(每对词都过一遍小网络)的简化——点积一次矩阵乘法就算完3

2. 复习并坐实:QKV 五步

第 01 章给过四步直觉,现在补上第 5 步(缩放),凑成正式的五步。 设定:每个词的表示已经拿着(第 04 章的 256 维),三份投影各自生成查询(我在找什么)、键(我是什么)、值(我能给什么)4

① 投影 每个词的表示 × 三张不同的权重矩阵 → Q、K、V 三份
② 打分 我的 Q · 每个 K(点积:对应位相乘再加总)
③ 缩放 分数 ÷ √64(键的维度)= 分数 ÷ 8
④ 掩码 不许看的格子填 −∞(下一节)
⑤ 加权 分数过 softmax 变成关注度,按关注度把所有 V 掺在一起

图说:第 ②③ 步合起来就是「缩放点积」这个名字;
第 ⑤ 步的输出形状和输入一样,所以块可以一层层往上码。

第 ③ 步为什么要除 √64? 原书给的是方差(数的波动幅度)论证:两个 64 维向量点积出来的分数, 维度越大波动越大;分数一大,softmax 就会变成「一家独大」的极端分布, 梯度(第 02 章)跟着消失。除以 √维度把波动拉回稳定区间5。 「这个除数看起来像拍脑袋,其实是让 softmax 别饱和」——原书还补了一句:不缩放会直接伤训练稳定性6

3. 主走查:一张 4×4 的掩码矩阵

掩码(mask,在打分之后、归一化之前,把「不许看」的格子的分数改成 −∞)是生成类模型的命门: 模型是逐词生成的,如果训练时能偷看「未来的词」,它学到的就是作弊7

拿一段 4 个词的句子走一遍(分数数值是为演示编的):

词序: w0 w1 w2 w3

│ w0 w1 w2 w3
────┼────────────────────
w0 │ 5 −∞ −∞ −∞ ← w0 是第一个词,只能看自己
w1 │ 3 4 −∞ −∞ ← w1 能看 w0 和自己
w2 │ 1 2 6 −∞
w3 │ 0 4 3 5 ← 最新的词,看得最全

图说:上三角全部 −∞(未来的词),下三角保留真实分数——这张「下三角」
就是原书说的因果掩码(亦称三角掩码)。softmax 之后 −∞ 格子的关注度恰好为 0,
未来的词对现在的词零贡献[^8]。

原书给的实现正好三行:造一张全 −∞ 的表,把下三角(含对角线)清零,加到分数上8。 书里引了个惊人的数字佐证「掩码必须对」:掩码做错,语言模型的困惑度(预测质量的一把尺,越低越好)会恶化 20-30%9

掩码家族还有两位常客,原书各给了一节10:

掩码用在哪干什么
填充掩码一批句子长短不齐时短句补的占位符不算分,免得污染平均值
滑动窗口掩码长文本(§4)只许看最近 w 个词,窗口外填 −∞

4. 三个省法之一:滑动窗口——只看最近的

因果掩码解决「不能看未来」,没解决「过去太长算不起」。滑动窗口注意力(SWA)的办法干脆: 每个词只看最近 w 个词,更早的一律 −∞11

成本账立刻改观:全量注意力 O(n²),滑窗 O(n·w)——长度 10,000、窗口 512 时, 原书算了只保留约 5% 的计算量(省 ~95%)12

代价呢?窗口外的远信息永久丢失——单层如此。但多层叠起来,信息可以一层层往外传: 第 1 层看 [t−512, t],第 2 层的 t−512 又能看 [t−1024, t−512]…… 原书还引了 GQA 那节之外的实证:注意力只算一小部分格子(这招叫稀疏:大多数格子直接不算)之后,效果保住 90-95%、内存省 80%,是这一族的常态口径13

5. 三个省法之二:GQA——键值共享

先立账本:推理时每个过去词的 K、V 都要存着(KV 缓存,第 04 章埋的显存账), KV 缓存的大小正比于「键值头的个数」。多头注意力(MHA,每路注意力各配全套 Q/K/V)在这里最贵。

GQA——正式名字是分组查询注意力——的刀法:查询头照旧多开(管多样性),键值头几个查询共用一套(管账单)14。 本书的配置(第 04 章)是 4 个查询头、2 个 KV 头——每套 K/V 供 2 个查询用, KV 缓存直接砍半;原书引的量产参照是 LLaMA-2 的 32 查询头/8 KV 头,缓存砍到四分之一15

原书还给了两个刻度,帮你在光谱上定位16:

方案键值头数位置
MHA=查询头数最贵,质量上限
GQA查询头数的 1/2 ~ 1/8量产主流
MQA只有 1 套最省,质量有损(GPT-3.5 Turbo 这类在线服务用)

代码上,共享的实现朴素得意外:把 K、V 沿一个新维度复制展开成和 Q 一样的头数, 然后照常算注意力——省的是,不是算前的形状17

6. 三个省法之三:sink——给多余的概率一个去处

这是一个反直觉的发现(原书引 Xiao 等人 2023):模型爱把大量注意力砸在开头几个词上, 哪怕它们毫无信息量——这个现象叫 attention sink(注意力沉底)18。 长文本生成时,开头 token 一旦被逐出缓存,质量立刻崩。

本书的处理不是硬压,而是顺水推舟:给每个头配一个可学习的常数分数(sink logits), 拼到分数矩阵的最后一列,一起做 softmax,然后把这列扔掉19:

分数 [4×4] → 拼一列 sink → [4×5] → softmax → 删掉 sink 列 → [4×4]

图说:sink 列在 softmax 时吸走「本来会硬塞给某个真实词」的多余概率,
删掉之后剩下的关注度加起来小于 1——这是故意的:多出来的概率有了去处,
真实词之间的相对关注度反而更干净[^21]。

7. 作者的判断与证据

说法书里的证据我们的标注
SDPA 是所有 transformer 的公共底座机制推导 + 「universal primitive」的定性20领域共识
掩码做错直接伤模型引困惑度恶化 20-30% 的研究9有出处;数字属二手转述
GQA 省缓存几乎不伤质量LLaMA-2 采用 32/8 的量产事实15厂商实践背书
sink 现象真实存在引 Xiao et al. 2023 的发现18有论文出处,原书如实标注

判断(我们的,不是书里的): 三个省法其实共用同一个思想—— 把「每个词看所有词」改成「按需分配」:滑窗按距离分配,GQA 按查询共享分配,sink 给用不掉的概率一个下水道。 原书把它们当三个独立技巧讲;抓住「按需分配」这条线,三个机制就能一起记住。 如果错,会错在: 如果某个省法(如 sink)的动机其实是纯数值稳定而非「分配」, 这个统一叙事就牵强——但作为帮助记住的框架,它不依赖归因完全正确。

8. 边界与局限

  • 原书自己列了掩码的代价:设计要懂领域、过度掩码会丢远处关键信息(书里引了文档问答实验:F1——兼顾查全与查准的评分——掉 5-10%)、 掩码让注意力矩阵更难解读21;
  • 滑窗窗口大小 w 是新的超参,任务不同最优值不同;本书取「偶数层开窗、奇数层全局」的折中(第 07 章);
  • 原书提的 FlashAttention(把注意力融进一个 GPU 核、不落盘中间矩阵)只给了一句定位22,展开在第 13 章;
  • 41:29 处「GPT-4 与 PaLM 坐拥万亿级参数」是市场传闻口径,不是可核的事实,引用时注意。

9. 可带走的

  1. SDPA 五步:投影→点积打分→除 √维度→掩码→softmax 加权;输出形状与输入相同,所以能层层堆;
  2. 除 √维度不是仪式:防 softmax 饱和、保梯度;
  3. 因果掩码 = 分数矩阵的上三角填 −∞;训练时偷看未来等于作弊;
  4. 注意力 O(n²) 是长文贵的一切根源;三个省法都在改「每个词看所有词」;
  5. 滑窗把成本压到 O(n·w),远信息靠层叠接力;
  6. GQA = 查询头多开、键值头共享,KV 缓存按共享倍数省;
  7. sink = 给 softmax 多余的概率一个虚拟去处,治「注意力都砸在开头」;
  8. 掩码是硬约束,做错的代价比做对省下的算力贵得多。

10. 原文地图

主题原书章原文位置
SDPA 出处与地位What Is Scaled Dot-Product Attention (SDPA)?text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:17(搜「Attention is All You Need」)
加性注意力的前身What Is Scaled Dot-Product Attention (SDPA)?text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:25(搜「Bahdanau」)
数据库/侦探类比What Is Scaled Dot-Product Attention (SDPA)?text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:33(搜「database query」) · text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:35(搜「river bank」)
O(n²) 成本What Is Scaled Dot-Product Attention (SDPA)?text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:139(搜「quadratic complexity」)
因果掩码定义Causal Maskingtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:57(搜「Causal masking」)
−∞ 实现与 tril 代码Causal Maskingtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:71(搜「large negative value」) · text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:77(搜「tril」)
掩码做错的代价Causal Maskingtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:85(搜「20–30%」)
填充掩码Padding Maskingtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:89(搜「Padding masking」)
滑窗省 95%Sparse Attention Maskstext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:139(搜「Sliding Window Attention」) · text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:141(搜「95%」)
稀疏化 90-95% 性能/80% 内存Sparse Attention Maskstext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:161(搜「90-95%」)
sdpa 函数签名与 sinkCustom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:309(搜「def sdpa」) · text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:325(搜「sink logits」)
GQA 复制展开 K/VCustom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:347(搜「unsqueeze(3)」)
因果与滑窗掩码取 maxCustom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:391(搜「torch.maximum」)
sink 拼列再切掉Custom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:403(搜「torch.cat」) · text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:405(搜「:-1」)
除 √D 防饱和Custom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:543(搜「sm_scale」)
sink=虚拟 keyCustom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:489(搜「virtual」)
剩余权重和小于 1 是故意的Custom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:611(搜「less than 1」)
T=4 window=2 例Custom Implementation of SDPA—Sliding Window and Grouped Query Attention for Our LLMtext/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:625(搜「T=4, sliding_window=2」)
GQA 定义与省账What Is Grouped Query Attention, Sliding Window Attention, and Sink Tokens?text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:317(搜「Grouped Query Attention (GQA)」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:327(搜「memory efficiency」)
sink 现象出处Limitations and Challengestext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:113(搜「Attention Sinks」)
「万亿参数」传闻口径What Is Scaled Dot-Product Attention (SDPA)?text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:29(搜「trillions of parameters」)

Footnotes

  1. 出处:「Challenges and Limitations of Attention」(Computational Cost 条)(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:139,搜「quadratic complexity」)。原文:全量自注意力对所有词对算分,复杂度随序列长度二次增长,长序列上即使现代 GPU 也吃力。

  2. 出处:「What Is Scaled Dot-Product Attention (SDPA)?」第 17 段(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:17,搜「Scaled dot-product attention」)。原文:SDPA 是 transformer 架构的基石,2017 年《Attention is All You Need》首次提出。

  3. 出处:「Historical Evolution and Contextual Foundations」(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:25,搜「Bahdanau」;:27,搜「dot product as a simpler」)。原文:2014 年加性注意力用小网络逐对算对齐分、计算密集;SDPA 改用点积并按维度平方根缩放,更简单高效。

  4. 出处:「Custom Implementation of SDPA」(函数签名段)(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:485,搜「queries」)。原文:Q 是「询问相关信息」的查询;K、V 形状一致,反映 GQA 中多个查询共享键值。

  5. 出处:「Custom Implementation of SDPA」(Step 2)(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:543,搜「sm_scale」)与「Softmax Scaling Factor」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:647,搜「Without scaling」)。原文:除以 √D 把点积方差归一,否则分数随维度增长、softmax 过尖、梯度受损;不缩放的高方差 logits 会把 softmax 推向接近二元的输出。

  6. 出处:「Training Stability」条(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:125,搜「Large dot products」)。原文:不缩放的大点积会破坏训练稳定性;另见 :647(搜「Without scaling」):高方差 logits 把 softmax 推向接近二元的输出、阻碍梯度流动。

  7. 出处:「Causal Masking」第 57 段(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:57,搜「Causal masking」)。原文:因果(自回归/三角)掩码保证位置 t 的查询只关注位置 ≤t 的键,防止模型获取未来信息、违反自回归性质。

  8. 出处:「Implementation Details」代码(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:77,搜「tril」)。原文:torch.tril 造下三角 1 矩阵,0 处填 −inf,softmax 前加到 logits 上。

  9. 出处:「Empirical Insights」(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:85,搜「20–30%」)。原文:研究(引 Vaswani et al. 2017)显示掩码不当会使语言模型困惑度上升 20-30%。 2

  10. 出处:「Padding Masking」(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:89,搜「Padding masking」)与「Sparse Attention Masks」(:139,搜「Sliding Window Attention」)。填充掩码把补位 token 的 logits 置 −∞ 保零权重;滑窗把窗口外置 −∞。

  11. 出处:「Sliding Window Attention (SWA)」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:329,搜「Sliding Window Attention (SWA)」)。原文:SWA 把注意力限制在每个 token 周围固定窗口,复杂度从二次降到线性;因果版本只看过去。

  12. 出处:「Sparse Attention Masks」(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:141,搜「95%」)。原文:k=512、n=10,000 时相比全量注意力省约 95% 计算(Longformer 一系)。

  13. 出处:「Empirical Insights」(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:161,搜「90-95%」)。原文:BigBird 的稀疏掩码在 NLP 基准上达到全量注意力 90-95% 的性能,同时省 80% 内存。

  14. 出处:「Grouped Query Attention (GQA)」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:317,搜「Grouped Query Attention (GQA)」;:327,搜「memory efficiency」)。原文:GQA 把查询分组、每组共享一套键值投影,把 KV 参数从 N_h 降到 N_kv,省显存且性能几乎不掉。

  15. 出处:「Parameter Extraction and GQA Configuration」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:525,搜「4x smaller」)。原文:32 查询头/8 KV 头 → KV 缓存缩为 1/4,对长序列与受限硬件部署关键(引 Ainslie et al. 2023)。 2

  16. 出处:「Comparison to Other Attention Mechanisms」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:835,搜「Standard MHA」;:837,搜「Multi-Query Attention (MQA)」)。原文:MHA 每头全套投影;MQA 只有一套 K/V,更极端;本书取 GQA 平衡表达力与效率。

  17. 出处:「Custom Implementation of SDPA」(Step 1)(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:507,搜「unsqueeze(3)」)。原文:K、V 沿新维度 expand 到与 Q 同头数——这是省显存的视图操作而非拷贝,组内每个查询用同一套键值。

  18. 出处:「Limitations and Challenges」(Attention Sinks 条)(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:113,搜「Attention Sinks」)。原文:模型可能过度关注开头 token(Xiao et al., 2023),使注意力分布偏斜。 2

  19. 出处:「Custom Implementation of SDPA」(Step 4)(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:597,搜「Sink logits」;:603,搜「virtual」)。原文:每头一个常数分数,cat 拼到分数最后一列,softmax 后用 [..., :-1] 切掉;sink 像「虚拟键」,吸收 softmax 的多余概率质量、帮助数值稳定。

  20. 出处:「The Enduring Legacy of SDPA」(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:303,搜「universal primitive」)。原文:SDPA 已成为 AI 的通用原语。

  21. 出处:「Challenges and Limitations of Masking Mechanisms」(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:247,搜「Over-Masking」;:249,搜「5-10%」)。原文:过度掩码限制模型表达力,文档级问答里 128 token 的窗口使 F1 掉 5-10%(引 Beltagy et al. 2020);掩码还让注意力权重更难解读。

  22. 出处:「Custom Implementation of SDPA」末段(text/41-ch06-6-scaled-dot-product-attention-core-sliding-wind.txt:621,搜「Flash Attention」)。原文:大 T 下每头 O(T²) 不可行,Flash Attention v2 用分块与重计算把显存降到 O(T);滑窗把复杂度降到 O(T·window),支撑 10 万级上下文。