LLM WIKI · 课程精读

LEARNING UNIT · 05

注意力机制与结构变体

比较 MHA、MQA、GQA、MLA、自注意力与交叉注意力,并解释缩放、掩码和残差的计算语义。

已整理章节
15 节
单元来源
14 条视频
总时长
38:52
状态
已发布
学习位置
5 / 20
01

主题讲解 · 02:06

可学习 Sink Logit 如何让注意力头抑制上下文

学习目标

  • 能写出普通 softmax 注意力与加入 sink logit 后的归一化公式。
  • 能解释为什么真实 token 的注意力权重和可以小于 1。
  • 能把 sink 机制等价理解为一个 value 为零的虚拟槽位。
  • 能说明每头独立 sink 参数如何形成上下文门控。
  • 能区分“抑制整个上下文混合”与“删除某个特定 token”。

前置与衔接

标准注意力头先计算某个 query 对所有可见 key 的 logits:

zj=qkjdk,j=1,,n.z_j=\frac{qk_j^\top}{\sqrt{d_k}}, \qquad j=1,\ldots,n.

普通 softmax 权重为

πj=ezji=1nezi.\pi_j=\frac{e^{z_j}}{\sum_{i=1}^{n}e^{z_i}}.

它们满足

j=1nπj=1.\sum_{j=1}^{n}\pi_j=1.

所以标准注意力输出是 value 的凸组合:

o=j=1nπjvj.o=\sum_{j=1}^{n}\pi_jv_j.

视频介绍的 attention sink 变体为每个头加入一个可学习 sink logit,让一部分 softmax 质量不再分给真实上下文 value。

图 1

视频所述机制为每个注意力头加入可学习的 sink logit,使一部分 softmax 质量流向不携带上下文 value 的虚拟槽。

原视频 · 00:00 ↗

核心讲解

1. 多头注意力先逐头计算

设有 hh 个注意力头。

每个头 rr 有自己的 Qr,Kr,VrQ_r,K_r,V_r,先独立计算

Zr=QrKrdk,Z_r=\frac{Q_rK_r^\top}{\sqrt{d_k}},

再用对应权重加权 VrV_r

图 2

多头注意力逐头计算分数与 value 加权;每个头可有独立 sink logit,从而学到不同的上下文抑制强度。

原视频 · 00:20 ↗

在拼接和输出投影之前,各头可采用不同的 sink logit srs_r

这意味着模型可以让某些头强烈依赖上下文,让另一些头把上下文贡献压低。

2. 普通 softmax 的真实权重和固定为 1

对某个 query 和某个头,记真实上下文 logits 为 z1,,znz_1,\ldots,z_n,并令

Z=i=1nezi.Z=\sum_{i=1}^{n}e^{z_i}.

普通权重是

πj=ezjZ.\pi_j=\frac{e^{z_j}}{Z}.
图 3

普通 softmax 只在真实上下文 logits 上归一化,因此每一行真实 token 的注意力权重和为 1。

原视频 · 00:40 ↗

不管这些 logits 整体多小,softmax 总会把全部概率质量重新分配给某些真实 token。

因此标准公式本身无法表达“这个头此刻不想读取任何上下文 value”;它至少要输出某个真实 value 的加权组合。

3. 在分母中加入虚拟 sink 质量

视频所述做法为该头增加一个可学习标量 ss

真实 token 权重改为

aj=ezjes+i=1nezi=ezjes+Z.a_j=\frac{e^{z_j}}{e^s+\sum_{i=1}^{n}e^{z_i}} =\frac{e^{z_j}}{e^s+Z}.
图 4

加入每头 sink logit s 后,softmax 分母多出 exp(s);真实 token 的分子不变,但它们获得的总质量下降。

原视频 · 01:00 ↗

若同时定义虚拟 sink 权重

asink=eses+Z,a_{\text{sink}}=\frac{e^s}{e^s+Z},

那么真实 token 与 sink 的总和仍然是 1。

只是 sink 不代表一个需要读取的真实上下文 token。

4. 为什么真实 token 权重和小于 1

真实权重之和为

j=1naj=ZZ+es.\sum_{j=1}^{n}a_j =\frac{Z}{Z+e^s}.

定义门控系数

g=ZZ+es.g=\frac{Z}{Z+e^s}.

只要 es>0e^s>0,就有

0<g<1.0<g<1.
图 5

真实上下文权重之和变为 Z/(Z+exp(s)),严格小于 1;s 越大,这个门控系数越接近 0。

原视频 · 01:20 ↗

所以视频所说的“注意力权重和小于 1”准确地说,是“所有真实上下文 token 的权重和小于 1”。

若把虚拟 sink 权重也算进去,扩展后的完整 softmax 权重和仍为 1。

5. 它等价于给普通注意力输出乘门

普通真实 token softmax 为

πj=ezjZ.\pi_j=\frac{e^{z_j}}{Z}.

加入 sink 后:

aj=ZZ+esezjZ=gπj.a_j =\frac{Z}{Z+e^s}\frac{e^{z_j}}{Z} =g\pi_j.

所以输出变成

osink=jajvj=gjπjvj=gostandard.o_{\text{sink}} =\sum_ja_jv_j =g\sum_j\pi_jv_j =g\,o_{\text{standard}}.
图 6

真实 V 的加权和整体被门控系数缩小;若虚拟 sink 不贡献 value,某个头可把上下文输出压到接近零。

原视频 · 01:40 ↗

在该公式约定下,sink 不改变真实 token 之间的相对比例:

aiaj=eziezj.\frac{a_i}{a_j}=\frac{e^{z_i}}{e^{z_j}}.

它改变的是整个真实上下文混合的总幅度。

6. 零 value 虚拟 token 的等价视角

可以把 sink 看成额外加入一个虚拟位置:

  • 它的 logit 是 ss
  • 它的 value 取零向量 vsink=0v_{\text{sink}}=0

完整输出是

o=j=1najvj+asink0.o=\sum_{j=1}^{n}a_jv_j +a_{\text{sink}}\cdot0.

softmax 仍完全归一化,但流向 sink 的概率质量对 value 加权和没有贡献。

这个视角能避免误以为 softmax 的数学性质被破坏。

7. sink logit 怎样控制抑制强度

门控系数是

g=ZZ+es=σ(logZs).g=\frac{Z}{Z+e^s} =\sigma(\log Z-s).

因此:

  • ss\to-\infty 时,es0e^s\to0g1g\to1,恢复普通注意力;
  • slogZs\gg\log Z 时,g0g\to0,真实上下文输出接近零;
  • 中间值产生连续的上下文门控。

注意 gg 不只由 ss 决定,也由当前 query 的真实 logits 总量 ZZ 决定。

一个固定的每头 ss 可以在不同 query 上产生不同实际门控强度。

8. “忽略上下文”不等于删除某个 token

这个机制直接缩放该头全部真实 value 的混合。

它不单独指定“忽略第一个 token”或“删除某一句话”。

若模型想在真实上下文内部偏向或避开某些 token,仍由 zjz_j 之间的相对差异决定。

所以更准确的说法是:某个头可以把“从真实上下文读取的总贡献”压到接近零。

其他头、残差连接和 FFN 仍可能携带上下文信息;单个头的 sink 不代表整个模型完全失去上下文。

9. 数值稳定实现

直接计算 ezje^{z_j}ese^s 可能溢出。

实现应像稳定 softmax 一样,先取

m=max(s,z1,,zn),m=\max(s,z_1,\ldots,z_n),

再计算

aj=ezjmesm+iezim.a_j= \frac{e^{z_j-m}} {e^{s-m}+\sum_i e^{z_i-m}}.

sink logit 必须参与行最大值和分母归约,才能保持数值等价。

10. 来源与版本边界

本课整理的是视频对其所称 DeepSeek V4 技术报告机制的讲解,并依据板书公式解释 sink logit。

当前课程素材没有附上该技术报告原文、精确张量 shape、参数共享方式或实现代码。

因此不能仅凭视频进一步断言:

  • sink 参数是否在所有层、所有头使用;
  • 是否还随 query 位置变化;
  • 是否存在额外正则、初始化或推理裁剪;
  • 具体版本是否采用完全相同的命名。

这些细节应以对应版本的正式报告与代码为准。

跟练与练习

原视频定位

编者练习

某个 query 的真实 logits 为 (1,2,3)(1,2,3),sink logit 为 4。写出真实 token 权重和 gg,并说明 sink logit 从 4 增大到 8 时会怎样。

查看参考答案

Z=e1+e2+e3Z=e^1+e^2+e^3,则真实权重和为 g=Z/(Z+e4)g=Z/(Z+e^4)。当 ss 从 4 增大到 8,分母中的 ese^s 大幅增加,gg 下降,真实上下文 value 混合被更强地压向零;真实 token 之间的相对权重比例仍不变。

编者练习 2

为什么说“加入 sink 后 softmax 权重和不再为 1”容易误导?

查看参考答案

若只统计真实 token,权重和确实是 Z/(Z+es)<1Z/(Z+e^s)<1。但把虚拟 sink 权重 es/(Z+es)e^s/(Z+e^s) 也算上,扩展 softmax 的总和仍为 1。变化的是一部分概率质量被分给了不贡献真实 value 的虚拟槽,而不是 softmax 不再归一化。

常见误区

  • 误区:attention sink 破坏了 softmax 权重和为 1。纠正:扩展后的真实 token 加虚拟 sink 仍归一化为 1。
  • 误区:sink 改变真实 token 之间的相对排序。纠正:在本课公式中它为所有真实权重乘同一门控 gg
  • 误区:sink 专门删除第一个 token。纠正:它抑制该头的整个真实上下文混合,不指定某个真实位置。
  • 误区:一个头忽略上下文等于整个模型忽略上下文。纠正:其他头与残差路径仍可携带信息。
  • 误区:sink logit 很大时输出严格为零。纠正:有限数值下通常只是接近零;精确效果还取决于真实 logits 总量。
  • 误区:视频板书足以确认所有版本实现细节。纠正:层范围、共享方式和代码细节需回到对应正式报告核验。

本课小结

  • 普通注意力把全部 softmax 质量分给真实上下文 token。
  • 每头加入 sink logit 后,分母多出 ese^s,真实权重和变为 g=Z/(Z+es)g=Z/(Z+e^s)
  • 机制等价于加入一个 logit 为 ss、value 为零的虚拟槽位。
  • 真实 value 混合整体乘以门控 gg,可被压到接近零。
  • sink 不改变真实 token 间相对比例,也不代表整个模型完全忽略上下文。
  • DeepSeek V4 的具体参数共享与实现范围应以对应版本正式材料为准。
02

主题讲解 · 02:16

为什么除以根号 d_k 等价于给注意力升温

学习目标

  • 能从 softmax 温度公式解释 1/dk1/\sqrt{d_k} 缩放的“升温”含义。
  • 能区分单头维度 dkd_k 与模型总宽度 dmodeld_{\text{model}}
  • 能在简化独立同分布假设下推导点积方差随 dkd_k 增长。
  • 能说明未缩放大 logits 为什么导致 softmax 饱和与梯度变弱。
  • 能区分数学缩放与 FlashAttention 等内核中的融合实现。

前置与衔接

缩放点积注意力的核心公式是

A=softmax(QKdk+M),A=\operatorname{softmax}\left( \frac{QK^\top}{\sqrt{d_k}}+M \right),

其中 MM 是 mask 或其他加性偏置。

softmax 温度形式是

softmax(zT).\operatorname{softmax}\left(\frac{z}{T}\right).

z=QKz=QK^\topT=dkT=\sqrt{d_k} 代入,两式完全同形。

图 1

缩放点积注意力把 QKᵀ 除以 sqrt(d_k) 后再做 softmax;相对未缩放 logits,这等价于使用温度 T=sqrt(d_k)。

原视频 · 00:00 ↗

dk>1d_k>1 时,dk>1\sqrt{d_k}>1,相对未缩放的 T=1T=1,这确实是更高的温度。

核心讲解

1. 温度怎样改变 softmax 尖锐度

对 logits ziz_i

pi(T)=ezi/Tjezj/T.p_i(T)=\frac{e^{z_i/T}}{\sum_je^{z_j/T}}.

两个候选的概率比为

pi(T)pj(T)=exp(zizjT).\frac{p_i(T)}{p_j(T)} =\exp\left(\frac{z_i-z_j}{T}\right).

因此:

  • T>1T>1 压缩 logit 差,分布更平;
  • 0<T<10<T<1 放大 logit 差,分布更尖;
  • 温度不改变 logits 排序,只改变相对概率差距。

“升温”在这里是数学类比,不是注意力模块真的具有物理温度。

2. d_k 是单个头的 key/query 维度

在常见多头注意力中,单个 query 头向量

qi,kjRdk.q_i,k_j\in\mathbb{R}^{d_k}.

dkd_k 是一次点积包含的特征数。

标准 MHA 常取

dmodel=hdk,d_{\text{model}}=h\,d_k,

但真正决定 qikjq_i^\top k_j 求和项数的是单头宽度 dkd_k,不是所有头拼接后的 dmodeld_{\text{model}}

在 GQA/MQA 等结构中,Q head 与 KV head 数量关系可以变化,但点积两端的单头维度仍必须兼容。

3. 为什么点积尺度随维度增长

把单个点积写成

s=qk=r=1dkqrkr.s=q^\top k=\sum_{r=1}^{d_k}q_rk_r.

做一个教学用初始化假设:

  • qr,krq_r,k_r 独立;
  • 均值为 0;
  • 方差为 1;
  • 各维乘积近似互不相关。

E[s]=0,\mathbb{E}[s]=0,
Var(s)=r=1dkVar(qrkr)dk.\operatorname{Var}(s) =\sum_{r=1}^{d_k}\operatorname{Var}(q_rk_r) \approx d_k.

所以标准差约为

Std(s)dk.\operatorname{Std}(s)\approx\sqrt{d_k}.
图 2

当 d_k 个乘积项累加使点积尺度增长时,softmax logits 差异会被放大,概率容易过早饱和。

原视频 · 01:00 ↗

除以 dk\sqrt{d_k} 后:

Var(sdk)1.\operatorname{Var}\left(\frac{s}{\sqrt{d_k}}\right)\approx1.

这个推导解释了初始化尺度选择,并不声称训练后的 Q/K 永远严格独立、方差永远等于 1。

4. 大 logits 为什么容易“赢家通吃”

softmax 只关心 logits 差。

以两个 logits 50 和 60 为例,用稳定 softmax 减去最大值:

softmax(50,60)=softmax(10,0).\operatorname{softmax}(50,60) =\operatorname{softmax}(-10,0).

较大项概率为

11+e100.99995.\frac{1}{1+e^{-10}}\approx0.99995.

较小项只有约 0.0000450.000045

图 3

大 logit 差会让注意力近似 one-hot,较小项及其梯度贡献显著减弱;这就是板书的赢家通吃直觉。

原视频 · 01:20 ↗

即使两个 token 都可能有用,value 加权和也几乎只留下较大 logit 对应的 VV

“赢家通吃”是接近 one-hot 的直觉说法,并不是权重在有限 logits 下数学上严格等于 0 或 1。

5. 饱和还会削弱梯度

softmax Jacobian 为

pizj=pi(δijpj).\frac{\partial p_i}{\partial z_j} =p_i(\delta_{ij}-p_j).

当某个 pi1p_i\approx1、其余概率接近 0 时,许多导数项也接近 0。

这会让训练很难通过小幅 logit 调整重新分配注意力。

1/dk1/\sqrt{d_k} 缩放的核心目的,是让 logits 的典型尺度不随头维度自动变大,降低无意义的早期 softmax 饱和。

它不保证注意力永远分散;模型仍可学习出真正尖锐的分布。

6. 为什么这叫“升温”

未经缩放的分数记为 S=QKS=QK^\top

标准注意力使用

softmax(Sdk).\operatorname{softmax}\left(\frac{S}{\sqrt{d_k}}\right).

对比温度公式可见

T=dk.T=\sqrt{d_k}.

相对 T=1T=1,更大的分母压缩分数差,因此称为给注意力“升温”。

图 4

除以 sqrt(d_k) 把典型点积尺度压回较稳定范围,使注意力不必因维度增长而自动变得更尖。

原视频 · 01:40 ↗

但这里比较的是“同一组未缩放 logits”与“缩放后 logits”。

训练后的模型参数会适应这套定义,不能把移除缩放后的模型当作同一个已训练函数来直接比较性能。

7. 与 CLIP 式温度的方向对照

视频用 CLIP 的对比学习做对照。

归一化图像/文本向量的余弦相似度位于 [1,1][-1,1],常再乘一个 logit scale,等价于除以温度:

zij=cos(ui,vj)T.z_{ij}=\frac{\cos(u_i,v_j)}{T}.

T<1T<1,相似度差被放大,分布变尖,即“降温”。

图 5

以 CLIP 式相似度为对照,小温度把有界相似度差异放大,使 softmax 更尖;具体温度可由实现固定或学习。

原视频 · 00:40 ↗

两者方向不同,是因为它们要校正的原始 logit 尺度不同:

  • 有界余弦相似度可能需要放大差异;
  • 高维未缩放点积可能需要压低典型尺度。

具体 CLIP 变体可能固定或学习 logit scale,不能把某个温度数值推广到所有版本。

8. 一个数值对比

仍取 logits (50,60)(50,60)

dk=100d_k=100,则 dk=10\sqrt{d_k}=10,缩放后为 (5,6)(5,6)

此时较大项概率为

e6e5+e6=11+e10.731.\frac{e^6}{e^5+e^6} =\frac{1}{1+e^{-1}} \approx0.731.

较小项仍保留约 0.269。

缩放不改变 60 大于 50 的排序,却显著缓和概率差。

9. mask 与缩放的次序

常见数学写法为

softmax(QKdk+M).\operatorname{softmax} \left(\frac{QK^\top}{\sqrt{d_k}}+M\right).

其中禁止位置的 Mij=M_{ij}=-\infty

也有实现先把 mask 加到未缩放分数,或用布尔 mask 在 softmax 内核中处理。

只要被禁止位置最终严格排除、允许位置缩放一致,语义可以等价。

有限精度下应避免把一个不够负的常数当作绝对 -\infty 后又因缩放改变屏蔽强度。

10. FlashAttention 中的算子融合

高性能注意力内核不会必须把完整 QKQK^\top、缩放结果和 softmax 矩阵分别写回显存。

它可以在分块计算中把以下步骤融合:

  • 点积累加;
  • 乘 scale 1/dk1/\sqrt{d_k}
  • 加 mask 或 bias;
  • 行最大值与在线 softmax 归约;
  • 与 V 的加权累加。
图 6

高性能注意力内核常把 scale、mask、稳定 softmax 等步骤融合,数学语义不变;具体融合边界取决于实现版本。

原视频 · 02:00 ↗

融合改变的是中间张量的存储和 kernel 调度,不改变缩放点积注意力的数学定义。

具体 scale 在点积前、点积后还是 softmax 内核中应用,取决于库、版本、数据类型和数值稳定策略。

11. 这项缩放不能保证什么

1/dk1/\sqrt{d_k} 不保证:

  • 每一行注意力都均匀;
  • 永远不会出现尖锐头;
  • 所有训练阶段 Q/K 都满足独立同分布;
  • 不再需要归一化、初始化或精度控制;
  • 任意自定义 attention score 都应机械使用同一比例。

它是标准点积注意力的尺度校准,让头维度变化时 logits 的典型量级更稳定。

跟练与练习

原视频定位

编者练习

dk=64d_k=64,未缩放的两个 logits 是 (16,24)(16,24)。求等效温度和缩放后的 logits,并比较两种情况下较大项的 softmax 概率。

查看参考答案

dk=8\sqrt{d_k}=8,等效温度 T=8T=8,缩放后 logits 为 (2,3)(2,3)。未缩放时差值为 8,较大项概率 1/(1+e8)0.99971/(1+e^{-8})\approx0.9997;缩放后差值为 1,较大项概率 1/(1+e1)0.7311/(1+e^{-1})\approx0.731。排序不变,但分布明显变平。

编者练习 2

若每个 qr,krq_r,k_r 的方差不是 1,而是都为 σ2\sigma^2,并保持独立零均值,点积方差约是多少?除以 dk\sqrt{d_k} 后还剩什么尺度?

查看参考答案

qrkrq_rk_r 的方差约为 σ4\sigma^4dkd_k 项相加得到 Var(qk)dkσ4\operatorname{Var}(q^\top k)\approx d_k\sigma^4。除以 dk\sqrt{d_k} 后方差约为 σ4\sigma^4。因此该缩放消掉的是随维度线性增长的部分,不会自动把任意输入方差都变成 1。

常见误区

  • 误区:dkd_k 是所有头拼接后的总宽度。纠正:它是单个点积头的 query/key 维度。
  • 误区:升温会改变最大 logit 的位置。纠正:正温度缩放不改变排序,只改变概率差距。
  • 误区:大点积必然源于 token 更相关。纠正:部分尺度增长只是求和维度增加造成的统计效应。
  • 误区:缩放后注意力一定均匀。纠正:模型仍可学习出有意义的大分数差。
  • 误区:初始化方差推导是训练后严格定律。纠正:独立同分布是假设,用于解释设计尺度。
  • 误区:FlashAttention 删除了除法。纠正:通常是把 scale 融合进内核,数学缩放仍然存在。

本课小结

  • softmax(QK/dk)\operatorname{softmax}(QK^\top/\sqrt{d_k}) 相对未缩放分数等价于温度 T=dkT=\sqrt{d_k}
  • 单头点积包含 dkd_k 项;在简化初始化假设下,其方差随 dkd_k 增长。
  • 除以 dk\sqrt{d_k} 把典型点积尺度稳定下来,减少无意义的 softmax 饱和。
  • 缩放不改变排序,也不禁止模型学习尖锐注意力。
  • CLIP 式小温度用于放大有界相似度差,与注意力缩放处理的原始尺度问题不同。
  • 高性能内核可以融合 scale、mask 和 softmax,但具体融合位置属于实现版本边界。
03

主题讲解 · 00:53

用 Shape 可视化多头注意力的拆头与拼头

学习目标

  • 能从 XX 开始写出 Q、K、V 投影、拆头、逐头注意力、拼头和 WOW_O 的完整 shape。
  • 能区分 dmodeld_{\text{model}}、头数 hh、单头维度 dkd_kdvd_v
  • 能解释“各头独立计算”发生在哪一段,以及 WOW_O 如何重新混合头信息。
  • 能识别 reshape、transpose、concat 分别改变什么张量语义。
  • 能检查缩放轴、softmax 轴和矩阵乘法的收缩维是否正确。

前置与衔接

设输入隐藏状态

XRB×N×dmodel,X\in\mathbb{R}^{B\times N\times d_{\text{model}}},

其中:

  • BB:batch size;
  • NN:序列长度;
  • dmodeld_{\text{model}}:模型隐藏宽度。

多头注意力不是把同一个完整注意力重复算 hh 遍。

它先把投影后的特征维组织成 hh 个较窄的子空间,每个头独立计算 attention,再拼回并做输出投影。

图 1

输入 X 的最后一维为 d_model;Q、K、V 投影后按 h 个头重塑,每头宽度 d_k 或 d_v。

原视频 · 00:00 ↗

核心讲解

1. 从 X 投影到 Q、K、V

为便于说明,先用标准 MHA,令

WQRdmodel×hdk,W_Q\in\mathbb{R}^{d_{\text{model}}\times hd_k},
WKRdmodel×hdk,W_K\in\mathbb{R}^{d_{\text{model}}\times hd_k},
WVRdmodel×hdv.W_V\in\mathbb{R}^{d_{\text{model}}\times hd_v}.

投影后:

Q=XWQRB×N×hdk,Q=XW_Q\in\mathbb{R}^{B\times N\times hd_k},
K=XWKRB×N×hdk,K=XW_K\in\mathbb{R}^{B\times N\times hd_k},
V=XWVRB×N×hdv.V=XW_V\in\mathbb{R}^{B\times N\times hd_v}.

很多经典配置取

dk=dv=dmodel/h,d_k=d_v=d_{\text{model}}/h,

于是投影总宽度仍等于 dmodeld_{\text{model}}

但这是常见配置,不是矩阵乘法本身强制的唯一选择。

2. 拆头是 reshape 加轴交换

QQ 的最后一维 hdkhd_k 拆成两个轴:

[B,N,hdk][B,N,h,dk].[B,N,hd_k] \rightarrow[B,N,h,d_k].

为了让 head 轴便于批量矩阵乘,再交换为

QRB×h×N×dk.Q\in\mathbb{R}^{B\times h\times N\times d_k}.

K,VK,V 同理。

图 2

Q、K、V 的投影维先按 head 轴拆分;实现通常是 reshape 加 transpose,而非必须复制成多份矩阵。

原视频 · 00:10 ↗

视频用蓝色与橙色把两个头叠放展示。

实现上通常只是 reshape/transpose 的视图语义;是否发生真实内存复制取决于后续 kernel 对 stride 和连续布局的要求。

3. 每个头独立计算 QKᵀ

rr 个头使用

QrRB×Nq×dk,Q_r\in\mathbb{R}^{B\times N_q\times d_k},
KrRB×Nk×dk.K_r\in\mathbb{R}^{B\times N_k\times d_k}.

得到分数

Sr=QrKrdkRB×Nq×Nk.S_r=\frac{Q_rK_r^\top}{\sqrt{d_k}} \in\mathbb{R}^{B\times N_q\times N_k}.

合并 head 轴后:

SRB×h×Nq×Nk.S\in\mathbb{R}^{B\times h\times N_q\times N_k}.
图 3

每个头用自身 Q_h、K_h、V_h 计算注意力;在输出拼接前,不同头的 QKᵀ 与 AV 不发生交叉相乘。

原视频 · 00:20 ↗

蓝色 Q 只与蓝色 K/V 配对,橙色头同理。

在这一阶段,不存在蓝色 Q 与橙色 K 的交叉内积。

4. 为什么缩放用 d_k 而不是 d_model

单个分数是

sij(r)=c=1dkQr,i,cKr,j,c.s_{ij}^{(r)} =\sum_{c=1}^{d_k}Q_{r,i,c}K_{r,j,c}.

求和项数是 dkd_k,因此缩放分母是

dk.\sqrt{d_k}.
图 4

每头分数矩阵为 Q_h K_hᵀ/sqrt(d_k),其序列轴 shape 为 n_q×n_k;d_k 是单头维度,不是 d_model。

原视频 · 00:30 ↗

dmodeld_{\text{model}} 是所有头拼接后的模型宽度,不能机械代替单头点积维度。

5. mask 与 softmax 的轴

加入 mask 后,对每个 batch、每个 head、每个 query 行,在 key 轴做 softmax:

Ab,r,i,:=softmax(Sb,r,i,:+Mb,r,i,:).A_{b,r,i,:} =\operatorname{softmax}(S_{b,r,i,:}+M_{b,r,i,:}).

所以

j=1NkAb,r,i,j=1\sum_{j=1}^{N_k}A_{b,r,i,j}=1

只对合法 key 成立。

softmax 不是沿 head 轴做;不同 head 的权重不需要彼此相加为 1。

mask 可通过广播扩展到 [B,h,Nq,Nk][B,h,N_q,N_k],但必须核对广播轴是否对应预期语义。

6. 每头用 A 加权 V

rr 个头的 value 为

VrRB×Nk×dv.V_r\in\mathbb{R}^{B\times N_k\times d_v}.

注意力输出

Or=ArVrRB×Nq×dv.O_r=A_rV_r \in\mathbb{R}^{B\times N_q\times d_v}.

矩阵乘的收缩轴是 NkN_k

[Nq,Nk]×[Nk,dv][Nq,dv].[N_q,N_k]\times[N_k,d_v] \rightarrow[N_q,d_v].

这就是视频所说的“AAVV 是加权平均”。

每个 query 行用自己的 key 权重,对 VV 的序列行做加权求和。

7. 拼头恢复总特征宽度

所有头输出堆叠为

ORB×h×Nq×dv.O\in\mathbb{R}^{B\times h\times N_q\times d_v}.

先交换回

[B,Nq,h,dv],[B,N_q,h,d_v],

再把最后两轴拼接:

OcatRB×Nq×hdv.O_{\text{cat}} \in\mathbb{R}^{B\times N_q\times hd_v}.
图 5

各头输出 O_h 在特征轴拼接回 h·d_v,再经 W_O 投影到 d_model;W_O 才允许不同头的信息线性混合。

原视频 · 00:40 ↗

hdv=dmodelhd_v=d_{\text{model}},拼接后宽度与输入 XX 相同。

如果二者不同,也可以通过输出投影映射回模型宽度。

8. W_O 让不同头重新混合

输出投影为

WORhdv×dmodel,W_O\in\mathbb{R}^{hd_v\times d_{\text{model}}},
Y=OcatWORB×Nq×dmodel.Y=O_{\text{cat}}W_O \in\mathbb{R}^{B\times N_q\times d_{\text{model}}}.

“各头互不干扰”只准确描述逐头计算 QrKrQ_rK_r^\topArVrA_rV_r 的阶段。

拼接后的 WOW_O 可以对所有头的特征做线性组合,因此最终输出通道会混合不同头的信息。

后续层的 Q/K/V 投影还会继续在这个混合表示上工作。

9. 两头例的 shape 台账

B=1,N=3,dmodel=8,h=2,dk=dv=4.B=1,\quad N=3,\quad d_{\text{model}}=8, \quad h=2,\quad d_k=d_v=4.

则:

阶段Shape
XX[1,3,8][1,3,8]
投影后 Q,K,VQ,K,V[1,3,8][1,3,8]
拆头后 Q,K,VQ,K,V[1,2,3,4][1,2,3,4]
分数 SS[1,2,3,3][1,2,3,3]
每头输出 OO[1,2,3,4][1,2,3,4]
拼头 OcatO_{\text{cat}}[1,3,8][1,3,8]
输出 YY[1,3,8][1,3,8]

相同的 shape 不代表相同语义。

例如拆头前的 [1,3,8][1,3,8] 与拼头后的 [1,3,8][1,3,8] 都是 8 维特征,但中间已经经过注意力汇聚。

10. 多头并不必然把总注意力计算乘 h

若固定 dmodeld_{\text{model}}dk=dmodel/hd_k=d_{\text{model}}/h,所有头的分数计算量约为

hNqNkdk=NqNkdmodel.h\cdot N_qN_kd_k =N_qN_kd_{\text{model}}.

头数增加时,每头变窄,总体主阶不因 hh 直接再乘一遍。

但实际速度还受 kernel shape、并行度、内存布局和 head 数影响,不能只凭大 O 断言耗时完全不变。

11. 与 GQA/MQA 的边界

本课视频展示标准 MHA:Q、K、V 都按同样数量的头拆分,并逐头配对。

GQA 让多个 query heads 共享一组 KV head;MQA 则让所有 query heads 共享更少的 K/V heads。

这时 Q 的 head 数与 K/V 的 head 数不同,需要额外的组映射或广播语义。

不能机械照搬“蓝头只对应蓝头”的一一配对图,但每个 query head 仍只会与分配给它的兼容 KV head 计算。

跟练与练习

原视频定位

编者练习

输入 XX 的 shape 为 [2,128,768][2,128,768],头数 h=12h=12,且 dk=dv=64d_k=d_v=64。写出拆头后 Q、分数矩阵 A、每头输出和拼头后的 shape。

查看参考答案

拆头后 QQ[2,12,128,64][2,12,128,64];self-attention 分数/权重 AA[2,12,128,128][2,12,128,128];每头输出整体张量仍为 [2,12,128,64][2,12,128,64];交换并拼头后为 [2,128,12×64]=[2,128,768][2,128,12\times64]=[2,128,768]

编者练习 2

为什么 softmax 应沿最后的 NkN_k 轴,而不是沿 head 轴?

查看参考答案

对固定 query 和固定 head,注意力要在所有可见 key 位置之间分配权重,所以归一化轴是 NkN_k。不同 head 表示不同子空间的独立注意力分布,不需要互相竞争总概率质量;沿 head 轴 softmax 会改变算法语义。

常见误区

  • 误区:拆头就是复制 XX 多份。纠正:先通过不同投影通道得到 Q/K/V,再 reshape 出 head 轴,通常不要求复制完整输入。
  • 误区:缩放使用 dmodel\sqrt{d_{\text{model}}}。纠正:单头点积收缩维是 dkd_k
  • 误区:softmax 在 head 轴归一化。纠正:它对每个 query 行沿 key 位置轴归一化。
  • 误区:各头从头到尾永不交互。纠正:逐头 attention 独立,但拼接后的 WOW_O 会混合头特征。
  • 误区:头数增加必然把总复杂度乘以头数。纠正:固定模型宽度时,每头通常相应变窄。
  • 误区:标准 MHA 图可以无改动解释 GQA/MQA。纠正:后两者需要 query head 到共享 KV head 的组映射。

本课小结

  • XX 先投影为总宽度 hdkhd_khdvhd_v 的 Q、K、V,再 reshape 出 head 轴。
  • 每个头独立计算 QrKr/dkQ_rK_r^\top/\sqrt{d_k},softmax 沿 key 轴归一化。
  • ArVrA_rV_r[Nq,Nk][N_q,N_k][Nk,dv][N_k,d_v] 收缩为 [Nq,dv][N_q,d_v]
  • 各头输出在特征轴拼接为 hdvhd_v,再由 WOW_O 映射回 dmodeld_{\text{model}}
  • 逐头 attention 阶段相互独立,输出投影之后可以混合。
  • shape 检查应同时核对 head 轴、序列轴、收缩轴和广播 mask 的语义。
04

主题讲解 · 03:01

GQA 如何在 MQA 与 MHA 之间调节 KV 共享

学习目标

  • 能用 Query 头数 hqh_q 与 KV 头数 hkvh_{kv} 统一描述 MHA、GQA 和 MQA。
  • 能写出 Query 头到 KV 头的分组映射,并计算每组共享规模。
  • 能解释“逻辑广播”为何不要求复制 KV Cache。
  • 能估算 GQA 相对 MHA 的 KV Cache 缩减比例。
  • 能说明共享 K/V 后,各 Query 头为何仍可能产生不同注意力。
  • 能准确理解“插值”只是 KV 头数上的离散架构过渡。

前置与衔接

标准多头注意力把投影后的 Q、K、V 都拆成多个头。

设 Query 头数为 hqh_q,KV 头数为 hkvh_{kv}

视频先回顾 MHA:每个 Query 头都有对应的 K/V 头。

图 1

MHA 中每个 Query 头配有各自的 K/V 头,头内独立计算后再拼接输出。

原视频 · 00:20 ↗

为了聚焦共享关系,先假设:

hqmodhkv=0.h_q \bmod h_{kv}=0.

定义每个 KV 头服务的 Query 头数

g=hqhkv.g=\frac{h_q}{h_{kv}}.

gg 就是 group size。

本课只比较注意力头组织与推理缓存。

模型层数、序列长度、单头宽度和 dtype 暂时保持相同。

核心讲解

1. 三种结构只差多少组 K/V

hqh_qhkvh_{kv} 可以统一三种结构:

结构KV 头数每个 KV 头服务的 Query 头数
MHAhkv=hqh_{kv}=h_qg=1g=1
GQA1<hkv<hq1<h_{kv}<h_q1<g<hq1<g<h_q
MQAhkv=1h_{kv}=1g=hqg=h_q

MHA 让每个 Query 头拥有独立 K/V。

MQA 把所有 Query 头收拢到同一组 K/V。

图 2

MQA 保留多个 Query 头,但只存一组物理 K/V,由全部 Query 头共享。

原视频 · 01:00 ↗

GQA 则在两者之间保留若干组 K/V。

因此它不是把 MHA 与 MQA 的输出做加权平均。

也不是在两套权重之间做连续插值。

“插值”指可通过离散改变 hkvh_{kv},获得不同共享强度。

2. Query 头如何找到所属 KV 头

若 Query 头编号为

r{0,1,,hq1},r\in\{0,1,\ldots,h_q-1\},

一种连续分组映射是

m(r)=rg.m(r)=\left\lfloor\frac{r}{g}\right\rfloor.

rr 个 Query 头使用第 m(r)m(r) 个 K/V 头:

Sr=QrKm(r)dk,S_r=\frac{Q_rK_{m(r)}^\top}{\sqrt{d_k}},
Ar=softmax(Sr+Mr),A_r=\operatorname{softmax}(S_r+M_r),
Or=ArVm(r).O_r=A_rV_{m(r)}.

例如 hq=8,hkv=2h_q=8,h_{kv}=2 时,g=4g=4

Query 头 0033 使用 KV 头 00

Query 头 4477 使用 KV 头 11

图 3

GQA 把 Query 头分组,每组共享一组 K/V;组间 K/V 仍可不同。

原视频 · 02:20 ↗

具体代码也可能采用交错编号。

关键不是编号顺序,而是每个 Query 头存在确定的 KV 组映射。

3. 广播是计算语义,不是缓存复制

视频用虚拟 K/V 头帮助画出多个逐头矩阵乘法。

图 4

共享 K/V 可在计算中做逻辑广播;虚拟头不要求把 KV Cache 复制成多份。

原视频 · 01:20 ↗

这些虚拟头不必在显存里真实复制。

实现可以让不同 Query 头通过索引、stride 或 kernel 内部映射读取同一块 K/V。

于是物理缓存仍只有 hkvh_{kv} 组。

如果先显式 repeat 再做普通 MHA,数学结果可以相同。

但那会物化重复张量,抵消 GQA 的内存与带宽优势。

高效 kernel 通常直接表达 grouped mapping。

4. 共享 K/V 不等于所有头相同

对同一组内的两个 Query 头 rrss,有

Km(r)=Km(s),Vm(r)=Vm(s).K_{m(r)}=K_{m(s)},\qquad V_{m(r)}=V_{m(s)}.

但一般仍有

QrQs.Q_r\ne Q_s.

因此

QrKm(r)QsKm(s).Q_rK_{m(r)}^\top \ne Q_sK_{m(s)}^\top.

softmax 后的 ArA_rAsA_s 也通常不同。

最终头输出 OrO_rOsO_s 可以不同。

图 5

即使 K/V 相同,不同 Query 头仍会得到不同的分数、注意力权重和头输出。

原视频 · 01:40 ↗

共享削减的是 K/V 表征数量,不是 Query 的多头多样性本身。

不过共享会减少 K/V 侧的容量,模型质量与效率之间仍存在权衡。

5. KV Cache 为什么按 h_kv 缩减

忽略额外元数据时,一层、一个样本的 K/V 元素量近似为

NKV=2Lhkvdh,N_{KV}=2Lh_{kv}d_h,

其中:

  • LL 是已缓存 token 数;
  • dhd_h 是单个 KV 头宽度;
  • 因子 22 来自 K 与 V。

若 dtype 每元素占 bb 字节,则

bytesKV=2Lhkvdhb.\text{bytes}_{KV}=2Lh_{kv}d_hb.

在其他条件相同的前提下,相对 MHA 的比例为

GQA cacheMHA cache=hkvhq.\frac{\text{GQA cache}}{\text{MHA cache}} =\frac{h_{kv}}{h_q}.

例如 hq=32,hkv=8h_q=32,h_{kv}=8,缓存约为 MHA 的 1/41/4

hkv=1h_{kv}=1,则退化为 MQA,约为 MHA 的 1/321/32

这里没有计入页表、对齐、量化尺度、RoPE 处理或实现临时缓冲。

6. “独立头”的边界

每个 Query 头会独立计算自己的 Sr,Ar,OrS_r,A_r,O_r

但所有头输出随后会拼接:

Ocat=Concat(O0,,Ohq1).O_{\text{cat}}=\operatorname{Concat}(O_0,\ldots,O_{h_q-1}).

再经过

Y=OcatWO.Y=O_{\text{cat}}W_O.

WOW_O 可以混合不同头的通道。

所以“头彼此独立”只适用于逐头 attention 阶段。

不能扩张成整个注意力层从头到尾完全没有跨头交互。

7. 把离散轴画完整

图 6

调节 KV 头数即可沿离散轴连接 MQA、GQA 与 MHA:1、介于两者之间、等于 Query 头数。

原视频 · 02:40 ↗

固定 hqh_q 后:

  • hkv=1h_{kv}=1:MQA;
  • 1<hkv<hq1<h_{kv}<h_q:GQA;
  • hkv=hqh_{kv}=h_q:MHA。

若要求等大小分组,hkvh_{kv} 还应整除 hqh_q

因此可选点通常是 hqh_q 的若干因数,而不是任意实数。

跟练与练习

原视频定位

编者练习

hq=16,hkv=4h_q=16,h_{kv}=4,使用连续分组。 求 group size,并写出 Query 头 0,3,4,11,150,3,4,11,15 对应的 KV 头编号。

查看参考答案

g=hq/hkv=4g=h_q/h_{kv}=4
映射为 m(r)=r/4m(r)=\lfloor r/4\rfloor,因此头 000\to0303\to0414\to111211\to215315\to3

编者练习 2

同一模型若从 hq=hkv=32h_q=h_{kv}=32 的 MHA 改为 hq=32,hkv=4h_q=32,h_{kv}=4 的 GQA,在层数、序列长度、头宽和 dtype 不变时,KV Cache 主体约缩到多少?

查看参考答案

比例为 hkv/hq=4/32=1/8h_{kv}/h_q=4/32=1/8
这是缓存张量主体的理论比例,不保证端到端显存也精确缩为 1/81/8,因为还有模型权重、激活、页表、对齐和临时工作区。

常见误区

  • 误区:GQA 对 MQA 与 MHA 的权重或输出做连续插值。纠正:它离散调节 KV 头数与共享组数。
  • 误区:逻辑广播会把 KV Cache 真正复制 gg 份。纠正:高效实现可通过映射直接复用同一存储。
  • 误区:共享 K/V 后所有注意力头完全相同。纠正:Query 头不同,分数、概率与输出通常不同。
  • 误区:GQA 的缓存一定精确等于 MHA 的 hkv/hqh_{kv}/h_q。纠正:该比值只描述同配置下 K/V 张量主体。
  • 误区:逐头 attention 独立意味着层输出不混合头。纠正:拼接后的 WOW_O 会混合通道。
  • 误区:任意 hq,hkvh_q,h_{kv} 都能等分。纠正:标准等大小分组要求 hqh_q 能被 hkvh_{kv} 整除。

本课小结

  • MHA、GQA、MQA 可统一为固定 hqh_q、改变 hkvh_{kv} 的离散架构族。
  • 每个 KV 头服务 g=hq/hkvg=h_q/h_{kv} 个 Query 头。
  • 逻辑广播不需要复制 KV Cache,是 GQA/MQA 节省内存和带宽的关键。
  • Query 头仍各自计算,因此共享 KV 不会自动让注意力与输出完全相同。
  • 同条件下,GQA 的 KV Cache 主体相对 MHA 约按 hkv/hqh_{kv}/h_q 缩减。
  • “头独立”止于逐头计算;拼接后的 WOW_O 可以重新混合所有头。
05

主题讲解 · 02:14

用五矩阵骨架理解注意力前向与反向传播

学习目标

  • 能写出单头 scaled dot-product attention 的完整前向链。
  • 能严格区分 softmax 前分数 SS 与 softmax 后概率 PP
  • 能从上游梯度 G=L/OG=\partial\mathcal L/\partial O 推导 dV,dP,dS,dQ,dKdV,dP,dS,dQ,dK
  • 能解释 softmax 反向为何不是简单逐元素传递。
  • 能使用 shape 检查矩阵乘法转置方向。
  • 能区分视频的五矩阵记忆骨架与完整反向实现。

前置与衔接

为突出矩阵关系,先讨论单头 self-attention。

Q,KRN×dk,Q,K\in\mathbb{R}^{N\times d_k},
VRN×dv.V\in\mathbb{R}^{N\times d_v}.

前向由两段矩阵乘法组成。

第一段用 Q、K 的内积形成分数。

第二段用归一化后的权重对 V 做加权平均。

图 1

前向先由 QK^T 形成每个 Query 对各 Key 的分数,再用 softmax 权重汇聚 V。

原视频 · 00:20 ↗

视频为了画面简洁,把 softmax 前后都记为 A。

课程稿改用 SSPP,避免反向传播时混淆。

核心讲解

1. 完整前向链

加入常见缩放与 mask:

S=QKdk+M,S=\frac{QK^\top}{\sqrt{d_k}}+M,
P=softmaxrow(S),P=\operatorname{softmax}_{\text{row}}(S),
O=PV.O=PV.

shape 为

S,PRN×N,S,P\in\mathbb{R}^{N\times N},
ORN×dv.O\in\mathbb{R}^{N\times d_v}.

softmax 对每个 Query 对应的一整行、沿 Key 轴归一化。

图 2

softmax 后的注意力概率 P 与 V 相乘得到输出 O;P 与原始分数 S 必须区分。

原视频 · 00:40 ↗

若是 cross-attention,只需把两个序列长度写成 Nq,NkN_q,N_k

2. 从输出梯度开始

设损失为 L\mathcal L,上游传回

G=LORN×dv.G=\frac{\partial\mathcal L}{\partial O} \in\mathbb{R}^{N\times d_v}.

对矩阵乘法 O=PVO=PV,有两条分支:

dP=GV,dP=GV^\top,
dV=PG.dV=P^\top G.

shape 检查:

[N,dv][dv,N]=[N,N],[N,d_v][d_v,N]=[N,N],
[N,N][N,dv]=[N,dv].[N,N][N,d_v]=[N,d_v].
图 3

给定上游梯度 G=dO,先有 dP=G V^T;同时完整反向还包含 dV=P^T G。

原视频 · 01:00 ↗

视频重点画出 dP=GVdP=GV^\top

完整反向不能漏掉同时累积到 V 的 dV=PGdV=P^\top G

3. softmax 反向是行内耦合

P=softmax(S)P=\operatorname{softmax}(S) 不是逐元素独立函数。

对某一行,若 u=dPiu=dP_ip=Pip=P_i,则

dSi=p(ujujpj).dS_i =p\odot\left(u-\sum_j u_jp_j\right).

矩阵写法为

dS=P(dProwsum(dPP)),dS =P\odot\left( dP-\operatorname{rowsum}(dP\odot P) \right),

其中 rowsum 的结果在该行广播。

这一步来自 softmax Jacobian:

pjsk=pj(δjkpk).\frac{\partial p_j}{\partial s_k} =p_j(\delta_{jk}-p_k).

因此不能把视频图中的“dAdA”不加区分地同时当成 dPdPdSdS

4. 从 dS 回到 Q 与 K

忽略不可训练的 mask,分数为

S=QKdk+M.S=\frac{QK^\top}{\sqrt{d_k}}+M.

于是

dQ=dSKdk,dQ=\frac{dS\,K}{\sqrt{d_k}},
dK=dSQdk.dK=\frac{dS^\top Q}{\sqrt{d_k}}.

shape 检查:

[N,N][N,dk]=[N,dk].[N,N][N,d_k]=[N,d_k].
图 4

dP 必须经过逐行 softmax 反向得到 dS,之后才有 dQ=dS K 与 dK=dS^T Q。

原视频 · 01:20 ↗

视频口播的 dQ=dAKdQ=dA\,K 对应的是把缩放和 softmax 中间步骤折叠后的骨架。

严格实现必须恢复 dPdSdP\rightarrow dS1/dk1/\sqrt{d_k}

5. causal mask 的梯度边界

若 causal mask 把未来位置设为 -\infty,对应概率为 0。

这些被遮挡位置不会参与有效 softmax 归一化。

实现中对 masked logits 的 dSdS 应为 0。

若用有限大负数近似 mask,需要关注低精度下的数值行为。

mask 本身通常不是可训练参数,所以不对 M 求梯度。

6. 五矩阵骨架能记什么

视频把 Q、K、注意力、V、O 排成五个矩阵。

图 5

Q、K、注意力、V、O 的五矩阵骨架便于记忆主路径,但不能替代完整梯度公式。

原视频 · 01:40 ↗

前向可以记为:

Q,KSP,Q,K\rightarrow S\rightarrow P,
P,VO.P,V\rightarrow O.

反向从 dOdO 向两条乘法输入分叉:

dOdP,dV,dO\rightarrow dP,dV,
dPdSdQ,dK.dP\rightarrow dS\rightarrow dQ,dK.

这个视觉骨架适合记链路与转置位置。

但它没有把 softmax、scale、mask 的细节全部画出。

也不能只靠五个矩阵判断所有梯度都已计算。

7. 前向与反向的视觉对称

图 6

前向按 Q→K→P→V→O 阅读,反向从 O 的上游梯度逆向追踪到 V、P、K、Q。

原视频 · 02:00 ↗

矩阵乘法 C=ABC=AB 的通用规则是

dA=dCB,dA=dC\,B^\top,
dB=AdC.dB=A^\top dC.

所以每遇到一次矩阵乘法,都应向左右两个输入各回传一条梯度。

“逆序阅读”是记忆提示,不是把前向式简单倒写。

真正的转置方向仍需由微分和 shape 决定。

8. 多头与 batch 版本

实际张量常为

Q,KRB×h×N×dk.Q,K\in\mathbb{R}^{B\times h\times N\times d_k}.

前述公式对每个 batch、每个 head 独立成立。

矩阵乘法只作用在最后两个轴。

之后各头输出拼接,再由 WOW_O 混合。

完整注意力层的反向还需经过 concat、WOW_O 和 Q/K/V 投影层。

9. 与 FlashAttention 的版本边界

标准公式可以物化 SSPP

FlashAttention 通过分块、在线 softmax 与反向重计算避免保存完整 N×NN\times N 矩阵。

它计算的是数学等价的梯度。

但内核执行顺序和保存的中间量不同。

因此“五矩阵图”描述数学依赖,不描述具体高效 kernel 的内存访问流程。

跟练与练习

原视频定位

编者练习

给定 O=PVO=PV,其中 PR4×6P\in\mathbb{R}^{4\times6}VR6×8V\in\mathbb{R}^{6\times8}G=dOR4×8G=dO\in\mathbb{R}^{4\times8}。 写出 dPdPdVdV 的公式和 shape。

查看参考答案

dP=GVdP=GV^\top,shape 为 [4,8][8,6]=[4,6][4,8][8,6]=[4,6]
dV=PGdV=P^\top G,shape 为 [6,4][4,8]=[6,8][6,4][4,8]=[6,8]
一次矩阵乘法反向必须同时向两个输入分支回传。

编者练习 2

为什么 dPdP 不能直接拿来计算 dQ=dPKdQ=dP\,K

查看参考答案

因为 P=softmax(S)P=\operatorname{softmax}(S)dPdP 是对 softmax 输出的梯度,必须先乘 softmax 的行内 Jacobian 得到 dSdS
随后才有 dQ=dSK/dkdQ=dS K/\sqrt{d_k}。跳过这一步会同时漏掉概率耦合项和缩放因子。

常见误区

  • 误区:视频中的 A 从头到尾代表同一个量。纠正:严谨推导应区分 logits SS 与概率 PP
  • 误区:dP=dSdP=dS。纠正:softmax 的 Jacobian 会让同一行各位置的梯度互相耦合。
  • 误区:从 dOdO 只需算 dPdP。纠正:矩阵乘法还同时产生 dV=PdOdV=P^\top dO
  • 误区:dQ=dSKdQ=dS K 已是完整公式。纠正:scaled attention 还要乘 1/dk1/\sqrt{d_k}
  • 误区:五矩阵骨架就是完整 backward graph。纠正:它省略 scale、mask、softmax 与部分分支。
  • 误区:FlashAttention 使用不同数学梯度。纠正:它主要改变分块、重计算和内存访问,不改变目标导数。

本课小结

  • 前向为 S=QK/dk+MS=QK^\top/\sqrt{d_k}+MP=softmax(S)P=\operatorname{softmax}(S)O=PVO=PV
  • dOdO 先产生 dP=dOVdP=dO\,V^\topdV=PdOdV=P^\top dO
  • dPdP 必须经过逐行 softmax 反向得到 dSdS
  • 再由 dQ=dSK/dkdQ=dS K/\sqrt{d_k}dK=dSQ/dkdK=dS^\top Q/\sqrt{d_k} 回到 Q、K。
  • 五矩阵图适合记主路径,但必须配合 shape、softmax 与 mask 边界。
  • 高效注意力内核可不物化完整 S/PS/P,却仍计算与这些公式等价的结果。
06

主题讲解 · 02:19

Attention Sink 与温度如何以不同方式改变注意力

学习目标

  • 能写出带温度 softmax 与带 sink logit softmax 的定义。
  • 能证明温度会改变真实 token 之间的概率比。
  • 能证明单独加入 sink 会等比缩放全部真实 token 权重。
  • 能区分“重新分配概率”与“把概率质量导向额外槽位”。
  • 能说明 sink value 为零时,对上下文输出形成何种门控。
  • 能识别视频中“注意力池”不是常规 Attention Pooling 层。

前置与衔接

给定某个 Query 对 nn 个真实 token 的 logits

z1,,zn,z_1,\ldots,z_n,

标准注意力权重为

πi=eziZ,\pi_i=\frac{e^{z_i}}{Z},

其中

Z=j=1nezj.Z=\sum_{j=1}^{n}e^{z_j}.

标准 softmax 满足

i=1nπi=1.\sum_{i=1}^{n}\pi_i=1.

视频比较两种改变注意力汇聚的机制。

图 1

温度在真实 token 之间重新分配概率;Attention Sink 则把一部分总质量导向额外槽位。

原视频 · 00:00 ↗

这里口播所称“注意力池”对应 attention sink 或额外虚拟接收槽。

它不是把序列聚成单一向量的常规 Attention Pooling 层。

核心讲解

1. 温度直接缩放 logits

加入温度 T>0T>0 后:

πi(T)=ezi/Tjezj/T.\pi_i(T) =\frac{e^{z_i/T}} {\sum_j e^{z_j/T}}.

温度改变的是 logits 之间的相对差异。

图 2

温度通过缩放 logits 改变注意力的尖锐或平坦程度,而不是把全部权重乘同一常数。

原视频 · 00:40 ↗

T>1T>1,所有差值都缩小:

zizjT<zizj\frac{z_i-z_j}{T} <z_i-z_j

对正差值成立。

分布趋向平坦。

0<T<10<T<1,差值被放大,分布趋向尖锐。

2. 温度会改变 token 之间的概率比

任意两个 token 的概率比为

πi(T)πj(T)=exp(zizjT).\frac{\pi_i(T)}{\pi_j(T)} =\exp\left(\frac{z_i-z_j}{T}\right).

只要 zizjz_i\ne z_j,改变 TT 就会改变这个比值。

因此温度不是把所有权重乘同一个常数。

T>1T>1 时,较高概率通常下降,较低概率通常上升。

图 3

当 T>1 时,logit 差异被压缩:较大的权重下降,较小的权重上升。

原视频 · 01:00 ↗

“通常”是针对非均匀分布的直观描述。

具体某一项如何变化还取决于它相对整体分布的位置。

3. 温度仍保留总概率为 1

无论 TT 取何正值:

iπi(T)=1.\sum_i\pi_i(T)=1.
图 4

仅改变温度后 softmax 仍归一化为 1,因此它是在 token 之间重新分配质量。

原视频 · 01:20 ↗

温度只是把固定总量在真实 token 之间重新分配。

它没有凭空移走真实 token 的总概率质量。

在极限下:

  • T0+T\rightarrow0^+ 时,概率趋向最大 logit 的 one-hot;
  • TT\rightarrow\infty 时,概率趋向均匀分布。

并列最大值等特殊情况需单独考虑。

4. Attention Sink 是给分母增加一项

加入一个 sink logit ss,把它与真实 logits 一起 softmax。

真实 token 权重变为

ai=eziZ+es.a_i =\frac{e^{z_i}}{Z+e^s}.

sink 自身权重为

asink=esZ+es.a_{\text{sink}} =\frac{e^s}{Z+e^s}.
图 5

加入 sink logit s 后,分母从 Z 变为 Z+e^s,真实 token 得到的总概率小于 1。

原视频 · 01:40 ↗

所有位置一起仍归一化:

iai+asink=1.\sum_i a_i+a_{\text{sink}}=1.

但真实 token 的总和降为

iai=ZZ+es<1.\sum_i a_i=\frac{Z}{Z+e^s}<1.

5. Sink 对真实 token 是等比缩放

将标准权重 πi=ezi/Z\pi_i=e^{z_i}/Z 代入:

ai=ZZ+esπi.a_i =\frac{Z}{Z+e^s}\pi_i.

定义门控系数

g=ZZ+es.g=\frac{Z}{Z+e^s}.

ai=gπi.a_i=g\pi_i.
图 6

固定原 logits 时,所有真实 token 权重共同乘上 Z/(Z+e^s),彼此比例保持不变。

原视频 · 02:00 ↗

因此任意两项的比例保持不变:

aiaj=πiπj=ezizj.\frac{a_i}{a_j} =\frac{\pi_i}{\pi_j} =e^{z_i-z_j}.

这正是它与温度最关键的区别。

6. Sink value 决定输出多了什么

若真实 value 为 viv_i,sink value 为 vsv_s,输出为

o=iaivi+asinkvs.o =\sum_i a_iv_i+a_{\text{sink}}v_s.

vs=0v_s=0,则

o=giπivi=gostandard.o =g\sum_i\pi_iv_i =g\,o_{\text{standard}}.

此时 sink 对真实上下文输出形成标量门控。

vs0v_s\ne0,输出还会加入 sink value 的贡献。

因此“等比下调真实权重”不自动等价于“输出只缩小”。

要看额外槽位的 value 如何定义。

7. 温度与 sink 可以同时存在

一种组合定义是把真实 logits 与 sink logit 一起除以同一温度:

ai(T)=ezi/Tjezj/T+es/T.a_i(T) =\frac{e^{z_i/T}} {\sum_j e^{z_j/T}+e^{s/T}}.

此时温度既改变真实 token 比例,也改变它们与 sink 的相对竞争。

另一种实现可能只对真实 logits 施加某个 scale,再另行设置 sink。

这两种参数化不完全相同。

比较实验时必须先写清温度作用到哪些 logits。

8. 与 scaled dot-product attention 的关系

标准注意力常用

zi=qkidk.z_i=\frac{q^\top k_i}{\sqrt{d_k}}.

dk\sqrt{d_k} 可视作控制 logit 尺度的固定因子。

但它的主要动机是稳定点积方差与 softmax 数值尺度。

可学习或人为设置的 attention temperature 是更一般的额外尺度。

不能因为二者都出现在分母,就忽略其设计目的与取值方式的差别。

9. “汇聚”不等于异常

尖锐注意力有时正是模型需要的选择行为。

平坦注意力也不自动表示更健康。

Attention Sink 常用于描述模型把概率集中到某些特殊位置的现象或机制。

其好坏取决于任务、训练方式、上下文长度与实现。

本课只比较两种数学作用,不把任何一种分布形态定性为必然优劣。

跟练与练习

原视频定位

编者练习

给定标准权重 π=(0.2,0.3,0.5)\pi=(0.2,0.3,0.5),加入 sink 后真实 token 总质量为 g=0.4g=0.4。 求三个真实权重与 sink 权重。

查看参考答案

真实权重等比缩放为 a=gπ=(0.08,0.12,0.20)a=g\pi=(0.08,0.12,0.20)
真实 token 总和为 0.40.4,因此 sink 权重为 1g=0.61-g=0.6
三个真实 token 的相对比例 2:3:52:3:5 保持不变。

编者练习 2

为什么更高温度不能用一个统一门控系数 gg 表示为 πi(T)=gπi\pi_i(T)=g\pi_i

查看参考答案

两次 softmax 都在真实 token 上归一化为 1。
若所有项都乘同一个 gg,求和要求 g=1g=1;但当 logits 非均匀且温度变化时,各 token 的概率比会改变,所以结果通常不可能与原分布完全相同。

常见误区

  • 误区:本课“注意力池”指常规 Attention Pooling。纠正:这里指 attention sink 或额外虚拟接收槽。
  • 误区:温度把所有权重等比缩放。纠正:温度改变 token 间概率比,同时总和仍为 1。
  • 误区:sink 会改变真实 token 之间的排序和比例。纠正:固定真实 logits 时,单独加 sink 对它们等比缩放。
  • 误区:真实 token 总和小于 1 违反 softmax。纠正:把 sink 一并计入后总和仍为 1。
  • 误区:sink value 为零与非零没有区别。纠正:非零 value 会向输出添加额外向量贡献。
  • 误区:更平坦或更尖锐的注意力必然更好。纠正:分布形态必须结合任务与训练目标评价。

本课小结

  • 温度使用 zi/Tz_i/T,会改变真实 token 之间的概率比。
  • 更高温度通常压低峰值、抬高尾部,但真实 token 概率总和仍为 1。
  • Attention Sink 在 softmax 分母中加入 ese^s,接收一部分总概率质量。
  • 单独加入 sink 时,所有真实 token 权重共同乘 g=Z/(Z+es)g=Z/(Z+e^s),彼此比例不变。
  • sink value 为零时,真实上下文输出被 gg 门控;非零时还会增加 sink 向量贡献。
  • 两种机制可以组合,但必须明确温度究竟作用到哪些 logits。
07

主题讲解 · 01:28

MQA 如何用一组 K/V 服务多个 Query 头

学习目标

  • 能从 MHA 基线准确指出 MQA 改了哪些投影 shape。
  • 能解释多个 Query 头如何共享一组物理 K/V。
  • 能区分逻辑广播、虚拟头和真实缓存副本。
  • 能说明共享 K/V 后各头输出为何仍通常不同。
  • 能估算 MQA 相对 MHA 的 KV Cache 缩减比例。
  • 能识别 MQA 的效率收益与表达容量边界。

前置与衔接

设输入

XRB×N×dmodel.X\in\mathbb{R}^{B\times N\times d_{\text{model}}}.

标准 MHA 使用 hqh_q 个 Query 头和同样数量的 K/V 头。

图 1

MHA 基线中,多个 Query 头分别配有自己的 K/V 头,并在头内独立计算。

原视频 · 00:00 ↗

若每个头宽为 dk,dvd_k,d_v,则常见投影为

Q=XWQRB×N×hqdk,Q=XW_Q\in\mathbb{R}^{B\times N\times h_qd_k},
K=XWKRB×N×hqdk,K=XW_K\in\mathbb{R}^{B\times N\times h_qd_k},
V=XWVRB×N×hqdv.V=XW_V\in\mathbb{R}^{B\times N\times h_qd_v}.

拆头后,各 Query 头与同编号 K/V 头计算。

MQA 保留多 Query 头,却把 K/V 的物理头数压到 1。

核心讲解

1. MQA 的投影 shape

MQA 中 Query 投影仍为

WQRdmodel×hqdk,W_Q\in\mathbb{R}^{d_{\text{model}}\times h_qd_k},
QRB×hq×Nq×dk.Q\in\mathbb{R}^{B\times h_q\times N_q\times d_k}.

但 K/V 只投影出一组:

WKRdmodel×dk,W_K\in\mathbb{R}^{d_{\text{model}}\times d_k},
WVRdmodel×dv,W_V\in\mathbb{R}^{d_{\text{model}}\times d_v},
KRB×1×Nk×dk,K\in\mathbb{R}^{B\times 1\times N_k\times d_k},
VRB×1×Nk×dv.V\in\mathbb{R}^{B\times 1\times N_k\times d_v}.
图 2

MQA 保留 h_q 个 Query 头,但 K/V 各只投影出一组物理头,供所有 Query 共享。

原视频 · 00:20 ↗

不能只说“把 K/V 切少”。

真正变化是 K/V 投影输出宽度从 hqdk,hqdvh_qd_k,h_qd_v 变为 dk,dvd_k,d_v

2. 每个 Query 头怎样计算

rr 个 Query 头使用同一个 K 与 V:

Sr=QrKdk,S_r=\frac{Q_rK^\top}{\sqrt{d_k}},
Ar=softmax(Sr+Mr),A_r=\operatorname{softmax}(S_r+M_r),
Or=ArV.O_r=A_rV.

所有 r{0,,hq1}r\in\{0,\ldots,h_q-1\} 都引用同一份 K/V。

实现可把 K/V 的 head 轴广播到 hqh_q

3. “虚拟头”不占额外 KV Cache

视频以虚线画出广播后的 K/V 头。

图 3

同一组 K/V 可按 Query 头索引逻辑广播;虚线头表示计算视图,不表示 KV Cache 副本。

原视频 · 00:40 ↗

这些虚拟头用于说明逐头计算的对齐关系。

它们不要求执行

Krepeat=repeat(K,hq).K_{\text{repeat}}=\operatorname{repeat}(K,h_q).

高效 kernel 可以让所有 Query 头直接读取原 K/V 存储。

因此显存里仍只保存一组 K 和一组 V。

若代码显式 repeat 并生成连续副本,数学上仍是 MQA。

但内存行为会失去 MQA 的主要优势。

4. 为什么注意力头仍然不同

共享 K/V 意味着

Kr=Ks=K,K_r=K_s=K,
Vr=Vs=V.V_r=V_s=V.

但 Query 投影保留不同头:

QrQs.Q_r\ne Q_s.

于是通常有

Ar=softmax(QrK/dk)As.A_r =\operatorname{softmax}(Q_rK^\top/\sqrt{d_k}) \ne A_s.

再用不同 ArA_r 对相同 V 加权,仍可得到不同 OrO_r

图 4

Query 头彼此不同,所以即便共享 K/V,各头的注意力矩阵和输出通常仍不同。

原视频 · 01:00 ↗

因此 MQA 没有把多头注意力退化成单个 Query 头。

它压缩的是 K/V 侧的头容量。

5. KV Cache 的节省

解码时,一层保存每个历史 token 的 K 与 V。

MHA 的缓存元素量近似为

2Lhqdh.2Lh_qd_h.

MQA 的缓存元素量近似为

2Ldh.2Ld_h.

在同样层数、长度、头宽与 dtype 下:

MQA cacheMHA cache1hq.\frac{\text{MQA cache}}{\text{MHA cache}} \approx\frac{1}{h_q}.

例如 hq=32h_q=32,K/V 张量主体约缩到 MHA 的 1/321/32

真实端到端显存还包含权重、激活、缓存管理元数据和临时工作区。

因此不能把 1/hq1/h_q 直接当作整机显存比例。

6. 带宽收益比算力主阶更关键

自回归 decode 每步只产生少量 Query,却要读取长上下文的 K/V。

此时计算往往受内存带宽限制。

MQA 减少 K/V 读取量,所以可改善 decode 吞吐。

它也减少 K/V 投影参数与写入缓存的数据量。

但不同硬件、batch、上下文长度和 kernel 实现会改变实际收益。

不能仅由理论缓存比例推出同倍数加速。

7. 多头输出仍会被 W_O 混合

每个 Query 头得到

OrRB×Nq×dv.O_r\in\mathbb{R}^{B\times N_q\times d_v}.

拼接后

Ocat=Concat(O0,,Ohq1).O_{\text{cat}} =\operatorname{Concat}(O_0,\ldots,O_{h_q-1}).

再经过

Y=OcatWO.Y=O_{\text{cat}}W_O.

视频所说“各头独立”指逐头 attention 的阶段。

WOW_O 会在线性输出层混合来自不同 Query 头的特征。

8. 与 GQA 的关系

MQA 是 hkv=1h_{kv}=1 的极端共享形式。

GQA 取

1<hkv<hq,1<h_{kv}<h_q,

让每组 Query 共享一组 K/V。

因此 GQA 可以在缓存开销与 K/V 侧表达容量之间提供更多离散选项。

实际选择由训练配置、质量目标和推理系统共同决定。

跟练与练习

原视频定位

编者练习

B=2,N=128,hq=8,dk=dv=64B=2,N=128,h_q=8,d_k=d_v=64。 写出 MQA 拆头后 Q、K、V、注意力矩阵 A 和全部头输出 O 的 shape。

查看参考答案

Q 为 [2,8,128,64][2,8,128,64];K 与 V 都只有一个物理头,分别为 [2,1,128,64][2,1,128,64];A 为 [2,8,128,128][2,8,128,128];全部头输出 O 为 [2,8,128,64][2,8,128,64]
K/V 的 head 轴可逻辑广播到 8,但不必物化为 8 份。

编者练习 2

若两个 Query 头共享相同 K/V,什么条件下它们的注意力权重一定相同?

查看参考答案

在 mask、缩放与数值路径也相同的前提下,只要两个 Query 头产生完全相同的逐位置 Q 向量,它们与同一 K 的分数就相同,softmax 权重也相同。
一般 MQA 使用不同 Query 投影,所以不能默认该条件成立。

常见误区

  • 误区:MQA 把 Q、K、V 都变成一个头。纠正:Q 仍保留多个头,只有物理 K/V 头压到一组。
  • 误区:广播必须复制 K/V。纠正:可通过索引或 stride 在 kernel 内直接共享同一存储。
  • 误区:共享 K/V 后所有注意力头相同。纠正:不同 Q 仍会生成不同分数和权重。
  • 误区:KV Cache 缩到 1/hq1/h_q,整机显存就缩到 1/hq1/h_q。纠正:该比例只适用于缓存张量主体。
  • 误区:理论带宽缩减会带来同倍数速度提升。纠正:实际速度还受 kernel、batch、并行度和其他算子影响。
  • 误区:各头从头到尾互不影响。纠正:拼头后的 WOW_O 会混合头特征。

本课小结

  • MQA 保留 hqh_q 个 Query 头,只保留一组物理 K 与 V。
  • 每个 Query 头独立使用同一 K/V 计算分数、softmax 与加权输出。
  • 虚拟广播是计算视图,不需要生成额外 KV Cache 副本。
  • Query 头不同,所以共享 K/V 后注意力权重和输出仍通常不同。
  • 同配置下,MQA 的 KV Cache 主体约为 MHA 的 1/hq1/h_q
  • MQA 主要改善 decode 的缓存容量与带宽,但可能牺牲部分 K/V 侧表达容量。
08

主题讲解 · 03:36

简化 MLA 如何用低维潜变量压缩 KV Cache

学习目标

  • 能画出简化 MLA 从 XXQ,cKV,K,V,OQ,c_{KV},K,V,O 的完整路径。
  • 能解释为什么 decode 时缓存低维 cKVc_{KV},而不是完整 K/V。
  • 能写出下投影与两条上投影的 shape。
  • 能区分“共享低维潜变量”与 MQA/GQA 的“共享 KV 头”。
  • 能说明低秩瓶颈为何不等价于无约束 MHA。
  • 能明确视频简化版省略的 RoPE、权重吸收与实现边界。

前置与衔接

设当前层输入为

XRB×N×dmodel.X\in\mathbb{R}^{B\times N\times d_{\text{model}}}.

传统 MHA 在自回归 decode 中缓存每个历史 token、每个 KV 头的 K 与 V。

MLA 的核心思路是先把 K/V 相关信息压入更窄的潜变量。

图 1

简化 MLA 从 X 生成瞬时 Q 和压缩潜变量 c_KV,再临时解压 K/V;图中省略 RoPE 与完整权重吸收。

原视频 · 00:00 ↗

本课遵循视频的简化教学模型:

  • 只讨论推理路径;
  • 暂时忽略 RoPE 的独立位置编码分量;
  • 显式还原临时 K/V,再按普通多头注意力计算;
  • 不展开完整 MLA 的权重吸收优化。

这些假设有助于理解缓存压缩,但不等同于 DeepSeek 完整实现细节。

核心讲解

1. Query 在当前 decode 步是瞬时量

当前 token 的 Query 可以直接写为

Q=XWQ,Q=XW_Q,

其中

WQRdmodel×hqdk.W_Q\in\mathbb{R}^{d_{\text{model}}\times h_qd_k}.

拆头后

QRB×hq×Nq×dk.Q\in\mathbb{R}^{B\times h_q\times N_q\times d_k}.
图 2

自回归推理中当前步 Q 用完即可丢弃,因此 Q 的下投影与上投影可合并为单个等效 W_Q。

原视频 · 01:00 ↗

若 Query 原本先下投影再上投影,且中间没有非线性或其他不可合并操作:

Q=XWDQWUQ=XWQ,Q=XW_{DQ}W_{UQ} =XW_Q,

其中

WQ=WDQWUQ.W_Q=W_{DQ}W_{UQ}.

decode 不需要把过去 token 的 Q 保存到 KV Cache。

所以视频把 Query 压缩路径合并成一个等效投影。

2. K/V 先下投影到低维潜变量

令压缩维度为 dcd_c,并取

dchkvdh.d_c\ll h_{kv}d_h.

下投影矩阵为

WDKVRdmodel×dc.W_{DKV} \in\mathbb{R}^{d_{\text{model}}\times d_c}.

压缩潜变量为

cKV=XWDKVRB×N×dc.c_{KV}=XW_{DKV} \in\mathbb{R}^{B\times N\times d_c}.
图 3

输入 X 经 W_DKV 下投影到较窄的 c_KV;跨 decode 步保存的是这份低维潜变量。

原视频 · 01:20 ↗

对于已经处理过的历史 token,缓存的是 cKVc_{KV}

新增 token 到来时,只需计算并追加它自己的潜变量行。

3. 从同一潜变量分支还原 K 与 V

简化版用两套不同上投影:

WUKRdc×hkvdk,W_{UK}\in\mathbb{R}^{d_c\times h_{kv}d_k},
WUVRdc×hkvdv.W_{UV}\in\mathbb{R}^{d_c\times h_{kv}d_v}.

得到

K=cKVWUK,K=c_{KV}W_{UK},
V=cKVWUV.V=c_{KV}W_{UV}.
图 4

c_KV 分别经 W_UK 与 W_UV 生成临时 K、V;两条不同上投影允许产生不同特征。

原视频 · 02:00 ↗

K 与 V 都源于同一潜变量,并不意味着二者相同。

WUKW_{UK}WUVW_{UV} 是不同参数,会提取不同特征。

同理,各 head 对应上投影矩阵的不同列块,也可以产生不同 K/V 头。

4. 缓存落点决定内存规模

传统 K/V 缓存每个 token 的元素量近似为

dtraditional=hkvdk+hkvdv.d_{\text{traditional}} =h_{kv}d_k+h_{kv}d_v.

dk=dv=dhd_k=d_v=d_h,则

dtraditional=2hkvdh.d_{\text{traditional}}=2h_{kv}d_h.

简化 MLA 的核心 latent 每 token 只需

dcd_c

个元素。

理想化主体比例为

MLA latent cachetraditional KV cachedc2hkvdh.\frac{\text{MLA latent cache}} {\text{traditional KV cache}} \approx \frac{d_c}{2h_{kv}d_h}.

完整 MLA 还可能缓存与位置编码有关的额外分量。

所以真实比例不能只看 dcd_c

还要计入层数、dtype、序列长度、对齐与缓存管理开销。

5. 解压后仍可拆成多个逻辑头

将 K、V 的最后一维拆成 head 轴:

KRB×hkv×Nk×dk,K\in\mathbb{R}^{B\times h_{kv}\times N_k\times d_k},
VRB×hkv×Nk×dv.V\in\mathbb{R}^{B\times h_{kv}\times N_k\times d_v}.
图 5

解压后的 Q、K、V 再按头拆分并做标准逐头注意力;共享潜变量不等于共享相同 KV 头。

原视频 · 02:40 ↗

简化版随后做标准注意力:

Sr=QrKrdk+Mr,S_r=\frac{Q_rK_r^\top}{\sqrt{d_k}}+M_r,
Ar=softmax(Sr),A_r=\operatorname{softmax}(S_r),
Or=ArVr.O_r=A_rV_r.

这里按 MHA 风格画成一一配对。

若具体配置的 hqh_qhkvh_{kv} 不同,还需要相应的组映射。

6. 共享潜变量不等于共享相同 KV 头

MQA/GQA 让多个 Query 头直接引用同一物理 K/V 头。

MLA 则让多个 K/V 特征共同由低维 cKVc_{KV} 生成。

对第 rr 个头,可以写成

Kr=cKVWUK(r),K_r=c_{KV}W_{UK}^{(r)},
Vr=cKVWUV(r).V_r=c_{KV}W_{UV}^{(r)}.

不同列块 W(r)W^{(r)} 可让头间表示不同。

因此“共享潜变量”不能直接等同于“所有头拿到相同 K/V”。

但所有头都受同一个 dcd_c 维瓶颈约束。

7. 低秩瓶颈带来表达边界

因为

K=XWDKVWUK,K=XW_{DKV}W_{UK},

合成投影

WK=WDKVWUKW_K=W_{DKV}W_{UK}

的秩至多为 dcd_c

V 路径同理。

如果 dcd_c 小于无约束投影所需秩,MLA 表达的是低秩结构。

所以不能因为解压后 shape 与普通 K/V 相同,就断言它与任意 MHA 完全等价。

它通过训练在压缩、效率与质量之间寻找合适折中。

8. 临时解压不一定是完整实现的执行方式

视频显式构造完整 K/V,便于说明。

完整 MLA 可利用矩阵乘法结合律,把某些上投影吸收到 Query 或输出侧。

例如不含位置分支时:

QK=Q(cKVWUK)=(QWUK)cKV.QK^\top =Q(c_{KV}W_{UK})^\top =(QW_{UK}^\top)c_{KV}^\top.

这说明有机会直接让变换后的 Q 与压缩 latent 计算。

类似地,V 的上投影可与后续路径重新组合。

实际 kernel 是否物化完整 K/V,取决于实现与位置编码结构。

9. RoPE 为什么使完整 MLA 更复杂

RoPE 对 Q/K 的特定维度施加与位置相关的旋转。

这种位置相关变换不能总像固定线性投影那样直接吸收。

完整 MLA 通常把内容相关分量与解耦的 RoPE 分量分开处理。

本视频板书明确省略这一分支。

因此简化图适合理解 latent cache,不适合直接照抄成 DeepSeek MLA 生产实现。

10. 拼头与输出投影

各头输出拼接为

Ocat=Concat(O0,,Ohq1).O_{\text{cat}} =\operatorname{Concat}(O_0,\ldots,O_{h_q-1}).

再经过

Y=OcatWO.Y=O_{\text{cat}}W_O.
图 6

各头输出拼接后仍需 W_O 投影回模型维度;本图只覆盖简化推理路径。

原视频 · 03:20 ↗

WOW_O 允许不同头的信息混合并映射回模型宽度。

“各头独立”仍只描述逐头 attention 计算阶段。

跟练与练习

原视频定位

编者练习

hkv=16,dk=dv=128,dc=512h_{kv}=16,d_k=d_v=128,d_c=512。 忽略 RoPE 额外缓存,计算每 token 的传统 K/V 元素量、MLA latent 元素量及理想化比例。

查看参考答案

传统缓存每 token 为 2hkvdh=2×16×128=40962h_{kv}d_h=2\times16\times128=4096 个元素。
MLA latent 为 dc=512d_c=512 个元素,理想化比例为 512/4096=1/8512/4096=1/8
真实实现还需加上位置相关缓存、对齐和管理元数据,不能据此直接断言总缓存精确缩为 1/81/8

编者练习 2

为什么 K=cKVWUKK=c_{KV}W_{UK} 的输出 shape 可以和普通 MHA 的 K 相同,却仍不能证明两者表达能力完全相同?

查看参考答案

因为 cKV=XWDKVc_{KV}=XW_{DKV},所以合成投影为 WK=WDKVWUKW_K=W_{DKV}W_{UK},其秩至多为 dcd_c
输出 shape 只说明维度相同,不说明线性映射的可达集合相同;较窄的 dcd_c 会施加低秩约束。

常见误区

  • 误区:MLA 缓存完整 K/V 后再压缩。纠正:核心是先缓存低维 cKVc_{KV},需要时才恢复或吸收计算。
  • 误区:Q 也必须跨 decode 步缓存。纠正:当前 Query 用于读取历史 K/V,通常用完即可丢弃。
  • 误区:K/V 共享同一 latent,所以所有头相同。纠正:不同上投影列块可生成不同头特征。
  • 误区:解压后 shape 与 MHA 相同,就具有无约束 MHA 的全部表达能力。纠正:合成投影受 dcd_c 的低秩瓶颈限制。
  • 误区:视频简图就是完整 DeepSeek MLA。纠正:它省略解耦 RoPE、权重吸收与具体 kernel。
  • 误区:缓存理论比例等于端到端显存比例。纠正:还需计入位置分量、其他层状态和系统开销。

本课小结

  • 简化 MLA 先计算瞬时 Q,并把 K/V 信息下投影为低维 cKVc_{KV}
  • decode 跨步缓存的是 cKVc_{KV},而不是传统的完整 K/V 头张量。
  • cKVc_{KV} 可经不同上投影生成临时 K/V,再按头计算注意力。
  • 共享 latent 不等于 MQA/GQA 的共享相同 KV 头,但会施加共同低秩瓶颈。
  • 完整实现可通过权重吸收避免显式物化部分 K/V;RoPE 分支使这一过程更复杂。
  • 视频图是理解缓存压缩的教学模型,不能代替完整 DeepSeek MLA 实现规范。
09

主题讲解 · 03:42

从 Decoder Query 到 Encoder Memory 的交叉注意力

学习目标

  • 能区分 decoder self-attention 与 cross-attention 的 Q/K/V 来源。
  • 能写出目标长度 NtN_t、源长度 NsN_s 下的完整 shape。
  • 能解释为什么 cross-attention 输出行数跟随 Query,而不是 encoder memory。
  • 能区分 causal mask 与 encoder padding mask。
  • 能说明 decoder 各层如何重复读取 encoder 最终输出。
  • 能识别视频所示经典 decoder block 与其他现代变体的边界。

前置与衔接

视频用图像描述任务举例。

图像 encoder 输出五个视觉 token,文本 decoder 当前有三个 token。

先看 decoder 内部的 masked self-attention。

图 1

decoder 自注意力的 Q、K、V 都由同一文本隐藏状态投影,三枚文本 token 形成 3×3 注意力。

原视频 · 00:40 ↗

若文本状态为

XtRB×Nt×dmodel,X_t\in\mathbb{R}^{B\times N_t\times d_{\text{model}}},

则 self-attention 的 Q、K、V 都由 XtX_t 投影。

在本例中 Nt=3N_t=3,所以全序列分数矩阵是 3×33\times3

cross-attention 的关键变化不是公式,而是信息来源。

核心讲解

1. Self-attention 的同源性

对第 rr 个头:

Qrself=XtWQ,rself,Q_r^{self}=X_tW_{Q,r}^{self},
Krself=XtWK,rself,K_r^{self}=X_tW_{K,r}^{self},
Vrself=XtWV,rself.V_r^{self}=X_tW_{V,r}^{self}.

因此

Srself=Qrself(Krself)dkRB×Nt×Nt.S_r^{self} =\frac{Q_r^{self}(K_r^{self})^\top}{\sqrt{d_k}} \in\mathbb{R}^{B\times N_t\times N_t}.

decoder 训练时还要加入上三角 causal mask。

ii 个文本位置不能读取 j>ij>i 的未来文本。

2. Cross-attention 的异源性

设 self-attention 后的 decoder 状态为

HtRB×Nt×dmodel,H_t\in\mathbb{R}^{B\times N_t\times d_{\text{model}}},

encoder 最终输出为

HsRB×Ns×denc.H_s\in\mathbb{R}^{B\times N_s\times d_{enc}}.

cross-attention 使用

Qrcross=HtWQ,rcross,Q_r^{cross}=H_tW_{Q,r}^{cross},
Krcross=HsWK,rcross,K_r^{cross}=H_sW_{K,r}^{cross},
Vrcross=HsWV,rcross.V_r^{cross}=H_sW_{V,r}^{cross}.
图 2

交叉注意力的 Q 来自 decoder 路径,K/V 来自 encoder 最终输出;两套投影参数与自注意力参数分开。

原视频 · 01:40 ↗

“交叉”指 Query 与 K/V 来自两个状态序列。

自注意力与交叉注意力的投影参数通常不共享。

3. 每个 decoder 层读同一份 encoder memory

经典 encoder-decoder Transformer 先完成 encoder。

其最终输出 HsH_s 作为 memory,供每个 decoder 层读取。

decoder 第 \ell 层通常有自己的

WQcross,(),WKcross,(),WVcross,().W_{Q}^{cross,(\ell)},W_{K}^{cross,(\ell)},W_{V}^{cross,(\ell)}.

所以输入 memory 相同,不代表各层投影后的 K/V 相同。

自回归生成时,某层的 encoder K/V 不随输出 token 改变。

实现可以在该层跨 decode 步缓存投影后的 encoder K/V。

4. Shape 由两侧 token 数共同决定

每个头的 shape 为

QrRB×Nt×dk,Q_r\in\mathbb{R}^{B\times N_t\times d_k},
KrRB×Ns×dk,K_r\in\mathbb{R}^{B\times N_s\times d_k},
VrRB×Ns×dv.V_r\in\mathbb{R}^{B\times N_s\times d_v}.

本例 Nt=3,Ns=5N_t=3,N_s=5

图 3

五个 encoder token 使 Kᵀ 有五列、V 有五行,Query 侧仍保持三个 decoder token。

原视频 · 02:20 ↗

因此

Sr=QrKrdkRB×3×5.S_r =\frac{Q_rK_r^\top}{\sqrt{d_k}} \in\mathbb{R}^{B\times3\times5}.

softmax 沿最后的五个 encoder 位置归一化。

5. “无 mask”准确说法是什么

视频把 cross-attention 标成“无掩码”。

准确含义是:不使用 decoder self-attention 的上三角 causal mask。

图 4

经典 encoder-decoder 交叉注意力不加 decoder 式上三角因果 mask,但仍可屏蔽 padding 或无效 encoder 位置。

原视频 · 02:40 ↗

每个文本 Query 通常都能读取全部有效 encoder token。

但仍可能存在:

  • encoder padding mask;
  • 图像 patch 有效性 mask;
  • 任务定义的稀疏或局部可见性 mask。

所以不能把“无 causal mask”写成“任何 mask 都没有”。

6. A×V 为什么回到 Query 行数

Ar=softmax(Sr+Ms)RB×3×5.A_r=\operatorname{softmax}(S_r+M_s) \in\mathbb{R}^{B\times3\times5}.

VrRB×5×dvV_r\in\mathbb{R}^{B\times5\times d_v}

相乘:

Or=ArVrRB×3×dv.O_r=A_rV_r \in\mathbb{R}^{B\times3\times d_v}.
图 5

3×5 的注意力矩阵与 5 行 V 相乘,收缩 encoder 长度 5,输出恢复为三个 Query 行。

原视频 · 03:00 ↗

中间维度 Ns=5N_s=5 被收缩。

每个文本 Query 得到一个由全部有效视觉 values 加权的向量。

输出行数始终等于 Query 数 NtN_t

7. 多头版本

加入 head 轴后:

ARB×h×Nt×Ns,A\in\mathbb{R}^{B\times h\times N_t\times N_s},
ORB×h×Nt×dv.O\in\mathbb{R}^{B\times h\times N_t\times d_v}.

各头输出拼接,再经 WOcrossW_O^{cross} 映射回模型宽度。

head 轴不会改变 Nt×NsN_t\times N_s 的序列二维语义。

8. 经典 decoder 层的数据流

图 6

视频所示经典 decoder 层依次经过 masked self-attention、cross-attention 与 FFN,并把三行输出传给下一层。

原视频 · 03:20 ↗

视频所示经典顺序是:

  1. masked self-attention;
  2. cross-attention;
  3. FFN;
  4. 交给下一 decoder 层。

每个子层实际还会配 residual 与 normalization。

Pre-Norm、Post-Norm、并行子层或门控结构会调整细节。

“FFN 每层只在最后一个”只描述该示意架构,不是所有模型的强制规律。

9. 训练与增量 decode

训练时三个目标 token 可以并行形成 Query。

causal mask 保证位置 ii 不读未来文本。

cross-attention 对 encoder memory 通常不需要因果限制。

增量 decode 时,当前 Query 长度常为 1。

这时 cross-attention 分数 shape 为

[B,h,1,Ns].[B,h,1,N_s].

encoder memory 仍然保持 NsN_s 个位置。

跟练与练习

原视频定位

编者练习

decoder 有 Nt=4N_t=4 个 Query token,encoder memory 有 Ns=7N_s=7 个 token,h=8,dk=dv=64h=8,d_k=d_v=64。 写出拆头后 Q、K、V、A、O 的 shape。

查看参考答案

Q 为 [B,8,4,64][B,8,4,64];K、V 均为 [B,8,7,64][B,8,7,64];A 为 [B,8,4,7][B,8,4,7];O 为 [B,8,4,64][B,8,4,64]
encoder 长度 7 在 A×VA\times V 中收缩,输出保留 Query 长度 4。

常见误区

  • 误区:cross-attention 的 Q、K、V 都来自 encoder。纠正:Q 来自 decoder,K/V 来自 encoder memory。
  • 误区:cross-attention 完全不使用 mask。纠正:通常没有 decoder causal triangle,但仍可有 padding/有效性 mask。
  • 误区:输出行数跟随 V。纠正:[Nt,Ns][Ns,dv][N_t,N_s][N_s,d_v] 输出 NtN_t 行。
  • 误区:所有 decoder 层复用同一套 cross-attention 投影。纠正:memory 可相同,各层参数通常独立。
  • 误区:FFN 必须紧跟每个 attention 子层。纠正:视频展示的是经典 decoder block 的特定顺序。
  • 误区:cross-attention 只能处理图像。纠正:encoder memory 可以来自文本、图像、音频或其他序列。

本课小结

  • self-attention 的 Q/K/V 同源,cross-attention 的 Q 与 K/V 异源。
  • cross-attention 分数一般为 [B,h,Nt,Ns][B,h,N_t,N_s]
  • softmax 沿 encoder/source token 轴归一化。
  • A×VA\times V 收缩 NsN_s,输出长度跟随 Query 的 NtN_t
  • 交叉注意力通常不使用 decoder 因果 mask,但仍可使用 source padding mask。
  • 经典 decoder 层的 self-attention、cross-attention、FFN 顺序需与具体模型版本一起理解。
10

主题讲解 · 03:45

为什么注意力分数按单头维度缩放

学习目标

  • 能解释缩放因子为什么使用 dkd_k 而不是总模型宽度。
  • 能在明确概率假设下推导点积分数的期望与方差。
  • 能证明除以 dk\sqrt{d_k} 会把理想方差从 dkd_k 恢复到 1。
  • 能说明使用 dmodel\sqrt{d_{model}} 在常见配置下会过度缩放多少。
  • 能把固定缩放与 softmax 饱和、梯度稳定联系起来。
  • 能区分教学分布假设与训练后真实 Q/K 的统计性质。

前置与衔接

scaled dot-product attention 使用

S=QKdk.S=\frac{QK^\top}{\sqrt{d_k}}.

这里 dkd_k 是单个 head 内 Query/Key 向量的宽度。

图 1

多头注意力在每个头内部做点积,收缩维是单头宽度 d_k,而不是全部头拼接宽度 d_model。

原视频 · 00:20 ↗

常见 MHA 配置满足

dmodel=hdk.d_{model}=h\,d_k.

但有些模型的 Q/K 总投影宽度不等于 dmodeld_{model}

因此应从实际点积的收缩维确定缩放,而不是机械读取模型隐藏宽度。

核心讲解

1. 为什么只看单头宽度

rr 个 head 的一个分数为

sij(r)=c=1dkqi,c(r)kj,c(r).s_{ij}^{(r)} =\sum_{c=1}^{d_k}q_{i,c}^{(r)}k_{j,c}^{(r)}.

求和发生在 dkd_k 个特征坐标上。

其他 heads 的坐标不会进入这个点积。

所以随机和的尺度由 dkd_k 决定。

dmodeld_{model} 只在常见配置中等于所有 head 宽度之和。

2. 推导需要哪些假设

为得到清晰公式,视频采用理想化假设。

对固定 i,ji,j,令

qc,kcq_c,k_c

在坐标 cc 上独立,并满足

E[qc]=E[kc]=0,\mathbb E[q_c]=\mathbb E[k_c]=0,
Var(qc)=Var(kc)=1.\operatorname{Var}(q_c) =\operatorname{Var}(k_c)=1.

不同坐标的乘积项也假定互不相关或独立。

图 2

推导使用期望线性、方差定义,以及独立变量乘积期望可分解等性质。

原视频 · 01:00 ↗

这些是假设,不是训练后张量逐项严格满足的事实。

其作用是说明为什么点积尺度随维度增长。

3. 未缩放点积的期望

s=c=1dkqckc.s=\sum_{c=1}^{d_k}q_ck_c.

期望的线性性给出

E[s]=c=1dkE[qckc].\mathbb E[s] =\sum_{c=1}^{d_k}\mathbb E[q_ck_c].

qc,kcq_c,k_c 独立时:

E[qckc]=E[qc]E[kc]=0.\mathbb E[q_ck_c] =\mathbb E[q_c]\mathbb E[k_c]=0.

因此

E[s]=0.\mathbb E[s]=0.

4. 点积方差为何逐项相加

在乘积项彼此独立的假设下:

Var(s)=c=1dkVar(qckc).\operatorname{Var}(s) =\sum_{c=1}^{d_k}\operatorname{Var}(q_ck_c).
图 3

在逐维乘积相互独立的近似下,q·k 的方差等于 d_k 个乘积项方差之和。

原视频 · 02:00 ↗

若坐标间存在协方差,完整公式还要加入交叉协方差项。

所以“方差可以直接相加”依赖独立或不相关条件。

5. 每个乘积项的方差

由方差定义:

Var(qckc)=E[qc2kc2]E[qckc]2.\operatorname{Var}(q_ck_c) =\mathbb E[q_c^2k_c^2] -\mathbb E[q_ck_c]^2.

第二项为 0。

独立性使第一项分解:

E[qc2kc2]=E[qc2]E[kc2].\mathbb E[q_c^2k_c^2] =\mathbb E[q_c^2]\mathbb E[k_c^2].

又因为

E[qc2]=Var(qc)+E[qc]2=1,\mathbb E[q_c^2] =\operatorname{Var}(q_c)+\mathbb E[q_c]^2=1,

K 同理,所以

Var(qckc)=1.\operatorname{Var}(q_ck_c)=1.
图 4

零均值、单位方差且独立时,每个 q_i k_i 的期望为 0、方差为 1。

原视频 · 02:40 ↗

6. 未缩放分数的方差增长为 d_k

dkd_k 个单位方差相加:

Var(s)=dk.\operatorname{Var}(s)=d_k.

标准差为

Std(s)=dk.\operatorname{Std}(s)=\sqrt{d_k}.
图 5

d_k 个单位方差项相加,使未缩放点积分数的方差增长到 d_k。

原视频 · 03:00 ↗

随着 head 变宽,典型 logit 幅度会变大。

softmax 更容易进入非常尖锐、接近饱和的区域。

7. 除以 sqrt(d_k) 后方差恢复

定义缩放分数

s~=sdk.\tilde s=\frac{s}{\sqrt{d_k}}.

E[s~]=0,\mathbb E[\tilde s]=0,

并由

Var(aX)=a2Var(X)\operatorname{Var}(aX)=a^2\operatorname{Var}(X)

得到

Var(s~)=1dkVar(s)=1.\operatorname{Var}(\tilde s) =\frac{1}{d_k}\operatorname{Var}(s) =1.
图 6

点积分数除以 sqrt(d_k) 后,方差再除以 d_k,在理想假设下恢复为 1。

原视频 · 03:20 ↗

这使初始化附近的 logit 尺度不随 head width 持续膨胀。

8. 一般方差版本

Var(qc)=σq2,Var(kc)=σk2,\operatorname{Var}(q_c)=\sigma_q^2, \qquad \operatorname{Var}(k_c)=\sigma_k^2,

并保持零均值与独立性,则

Var(qckc)=σq2σk2,\operatorname{Var}(q_ck_c) =\sigma_q^2\sigma_k^2,
Var(s)=dkσq2σk2.\operatorname{Var}(s) =d_k\sigma_q^2\sigma_k^2.

除以 dk\sqrt{d_k} 后仍剩

σq2σk2.\sigma_q^2\sigma_k^2.

所以固定因子消除的是维度增长,不会神奇地把任意实际分布都变成单位方差。

9. 为什么不是 sqrt(d_model)

在常见关系

dmodel=hdkd_{model}=h d_k

下,如果除以 dmodel\sqrt{d_{model}}

Var(sdmodel)=dkhdk=1h.\operatorname{Var} \left(\frac{s}{\sqrt{d_{model}}}\right) =\frac{d_k}{h d_k} =\frac1h.

相对目标尺度,它额外缩小了 h\sqrt h

softmax 分布会更平坦。

但“过度平坦”是相对该初始化推导而言。

训练可调整投影尺度,不能把一次方差分析写成所有模型行为的绝对定律。

10. 与温度的关系

softmax 温度形式为

pj(T)=esj/Tmesm/T.p_j(T)=\frac{e^{s_j/T}}{\sum_m e^{s_m/T}}.

从代数上看,除以 dk\sqrt{d_k} 相当于相对未缩放 logits 使用固定温度。

其首要设计动机是补偿点积随维度增长的尺度。

它不是可学习温度,也不保证注意力一定均匀。

模型仍可通过 Q/K 方向与范数形成尖锐注意力。

11. 数值与梯度直觉

logits 绝对差过大时,softmax 接近 one-hot。

许多非最大位置的梯度会非常小。

缩放让初始化附近更少落入饱和区,有利于优化稳定。

现代 fused attention 还会使用减最大值等数值稳定技巧。

这些技巧解决溢出;1/dk1/\sqrt{d_k} 解决的是统计尺度增长,两者作用不同。

跟练与练习

原视频定位

编者练习

dk=64d_k=64,各维 qc,kcq_c,k_c 独立、零均值,方差均为 2。 求未缩放点积与除以 64\sqrt{64} 后的方差。

查看参考答案

每个乘积项方差为 2×2=42\times2=4
未缩放点积方差为 64×4=25664\times4=256
除以 64=8\sqrt{64}=8 后,方差除以 82=648^2=64,得到 4。
缩放消除了维度因子,但不会把非单位输入方差自动改成 1。

常见误区

  • 误区:Q/K 投影总宽度永远等于 dmodeld_{model}。纠正:这是常见配置,不是所有架构的硬约束。
  • 误区:训练后 Q/K 元素严格 i.i.d. 标准正态。纠正:它只是尺度推导的近似假设。
  • 误区:和的方差总等于方差之和。纠正:需要独立或交叉协方差为零。
  • 误区:除以 dk\sqrt{d_k} 保证注意力均匀。纠正:它只稳定典型 logit 尺度。
  • 误区:用 dmodel\sqrt{d_{model}} 只是符号不同。纠正:常见配置下会额外缩小 h\sqrt h
  • 误区:减最大值与 1/dk1/\sqrt{d_k} 是同一技巧。纠正:前者防溢出,后者控制统计尺度。

本课小结

  • 注意力分数在单个 head 内计算,因此缩放维度是点积收缩轴 dkd_k
  • 在独立、零均值、单位方差假设下,未缩放点积期望为 0、方差为 dkd_k
  • 除以 dk\sqrt{d_k} 会使方差除以 dkd_k,恢复为 1。
  • dmodel\sqrt{d_{model}}dmodel=hdkd_{model}=h d_k 时会把方差压到 1/h1/h
  • 固定缩放降低初始化附近 softmax 饱和风险,但不决定最终注意力一定尖锐或平坦。
  • 实际解释必须区分理想统计推导、训练后分布与具体内核数值稳定策略。
11

主题讲解 · 03:51

Attention Residual 如何按 Token 混合历史层输出

学习目标

  • 能区分普通 residual 与视频所示 Attention Residual。
  • 能解释为什么候选 value 应取未再次混合的 clean branch 输出。
  • 能写出 residual query、RMSNorm key、value 与 softmax 权重。
  • 能说明同一层的权重为何仍可逐 token 变化。
  • 能检查候选数量、权重 shape 与输出 shape。
  • 能把视频示例权重与具体实现参数边界分开。

前置与衔接

普通 Transformer 子层常写为

y=x+F(x).y=x+F(x).

输入 xx 与当前子层输出 F(x)F(x) 都以系数 1 相加。

图 1

普通 residual 只把当前子层输入与输出相加;视频示意的 Attention Residual 可从多条历史 clean branch 动态混合。

原视频 · 00:00 ↗

视频所示 Attention Residual 不只读取紧邻输入。

它把多个历史 branch 输出视为候选 values,再动态加权。

本课按视频中的 Kimi 架构示意讲解。

不同论文或模型若也使用“attention residual”一词,具体公式未必完全相同。

核心讲解

1. 历史候选是什么

对固定 token ii,设到当前混合点共有 mm 条候选分支:

vi,1,vi,2,,vi,mRd.v_{i,1},v_{i,2},\ldots,v_{i,m} \in\mathbb{R}^{d}.

它们可来自词嵌入、早期 attention 直接输出、FFN 直接输出与当前子层直接输出。

视频第二层示例含四条候选。

图 2

第二层示例把词嵌入、早期 attention 输出、FFN 输出和当前 attention 输出按四个系数组合。

原视频 · 02:00 ↗

示例中的 0.1,0.2,0.1,0.60.1,0.2,0.1,0.6 只是一次可视化结果。

它们不是写死在网络中的固定常数。

2. 为什么从 clean branch 取值

图中弧线从每个子层箭头末端的小圆点出发。

图 3

弧线从子层直接输出的小圆点取值,而不是从已经混合过的 residual state 再递归取值。

原视频 · 02:20 ↗

这表示 value 使用该子层的直接输出。

若反复把已经聚合过的 residual state 再作为候选,早期信息会经过多层嵌套混合。

clean branch 让每条历史贡献有一条直接到当前混合点的路径。

视频把它解释为减轻早期信息被连环稀释的风险。

这是设计动机,不是对任意参数、任意深度的无条件数值保证。

3. Residual query

对第 \ell 个混合点,视频使用一个可学习向量

qres()Rd.q_{res}^{(\ell)}\in\mathbb{R}^{d}.

该 query 在同一层可对所有 token 共享。

它表达“当前层希望从哪类历史分支读取信息”。

不同层拥有不同 query,所以读取偏好可随深度变化。

4. 历史向量同时充当 key 与 value

视频没有额外的 K/V 投影矩阵。

对 token ii 的第 jj 条候选:

vi,j=clean branch output,v_{i,j}=\text{clean branch output},
ki,j=RMSNorm(vi,j).k_{i,j}=\operatorname{RMSNorm}(v_{i,j}).
图 4

每个混合位置使用一个可学习 residual query;历史分支本身作 value,并经 RMSNorm 后作为 key 参与匹配。

原视频 · 03:00 ↗

RMSNorm 用于匹配路径,value 保留原候选向量。

这样 key 的尺度更稳定,而输出仍是历史 features 的加权和。

具体实现是否还有 scale、bias 或门控,应以对应模型代码为准。

5. Softmax 生成分支权重

一种与视频一致的抽象写法是

si,j()=(qres())ki,j.s_{i,j}^{(\ell)} =(q_{res}^{(\ell)})^\top k_{i,j}.

对候选分支轴做 softmax:

αi,j()=esi,j()r=1mesi,r().\alpha_{i,j}^{(\ell)} =\frac{e^{s_{i,j}^{(\ell)}}} {\sum_{r=1}^{m}e^{s_{i,r}^{(\ell)}}}.
图 5

residual query 与同一 token 的历史 keys 匹配,softmax 得到对各 clean branch 的归一化权重。

原视频 · 03:20 ↗

所以

αi,j()0,j=1mαi,j()=1.\alpha_{i,j}^{(\ell)}\ge0, \qquad \sum_{j=1}^{m}\alpha_{i,j}^{(\ell)}=1.

当前输出为

yi()=j=1mαi,j()vi,j.y_i^{(\ell)} =\sum_{j=1}^{m}\alpha_{i,j}^{(\ell)}v_{i,j}.

6. 为什么权重逐 token 变化

qres()q_{res}^{(\ell)} 可以在 token 之间共享。

vi,jv_{i,j} 来自 token ii 自己的历史表示。

所以 ki,jk_{i,j} 也随 ii 改变。

因此

αi,:()αi,:()\alpha_{i,:}^{(\ell)} \ne \alpha_{i',:}^{(\ell)}

通常成立。

图 6

query 可在层内共享,但不同 token 的历史 key/value 不同,因此混合权重可逐 token 变化。

原视频 · 03:40 ↗

这里的动态轴是“历史分支”,不是序列中其他 token。

每个 token 对自己的历史层输出做混合。

7. Shape 台账

若序列长度为 NN,候选数为 mm,隐藏宽度为 dd

VhistRN×m×d,V_{hist}\in\mathbb{R}^{N\times m\times d},
KhistRN×m×d,K_{hist}\in\mathbb{R}^{N\times m\times d},
qres()Rd,q_{res}^{(\ell)}\in\mathbb{R}^{d},
S,αRN×m,S,\alpha\in\mathbb{R}^{N\times m},
YRN×d.Y\in\mathbb{R}^{N\times d}.

softmax 沿 mm 个历史候选归一化。

不能误沿 token 轴 NN 做归一化,否则会变成跨 token 混合。

8. 与普通 token attention 的差别

普通 self-attention 对固定 Query 在序列位置轴选择 K/V。

Attention Residual 对固定 token 在历史 branch/深度轴选择 K/V。

二者都使用 query-key 匹配与 softmax。

但被聚合的轴不同。

因此它不会自动替代当前层的 token self-attention。

9. 路径数与成本

层数增加时,可用历史候选数量可能增长。

显式保留所有 clean branch 会带来额外激活存储、读取和匹配成本。

实际系统可能限制候选窗口、分组、压缩或优化 kernel。

视频只解释语义,不给出完整复杂度与工程实现。

不能从示意连线数量直接推断生产模型的精确显存。

跟练与练习

原视频定位

编者练习

某 token 有三个历史 values,匹配 logits 为 [0,ln2,0][0,\ln2,0]。 求 softmax 权重,并写出输出加权式。

查看参考答案

指数为 [1,2,1][1,2,1],总和为 4,因此权重为 [1/4,1/2,1/4][1/4,1/2,1/4]
输出为
y=14v1+12v2+14v3.y=\tfrac14v_1+\tfrac12v_2+\tfrac14v_3.

常见误区

  • 误区:Attention Residual 只是给普通 skip connection 乘一个固定常数。纠正:视频中权重由 query-key 匹配动态产生。
  • 误区:权重在整批 token 上相同。纠正:query 可共享,但每个 token 的历史 keys 不同。
  • 误区:softmax 沿 token 轴。纠正:这里沿同一 token 的历史候选轴归一化。
  • 误区:value 应使用已经多次残差相加的 state。纠正:图中刻意使用 clean branch 直接输出。
  • 误区:示例 0.1,0.2,0.1,0.60.1,0.2,0.1,0.6 是模型固定参数。纠正:它们只是一次可视化权重。
  • 误区:clean path 一定杜绝信息衰减。纠正:它提供更直接路径,但效果仍由训练与参数决定。

本课小结

  • 普通 residual 混合紧邻输入与当前子层输出,Attention Residual 可读取多个历史 branch。
  • 候选 values 来自未再次混合的 clean branch 输出。
  • 层级 query 与每个 token 的 RMSNorm history keys 匹配,softmax 生成分支权重。
  • 同一层 query 可共享,但 keys 随 token 变化,所以权重逐 token 动态分配。
  • 输出在历史分支轴做加权和,shape 保持 [N,d][N,d]
  • 具体公式、候选范围与工程成本必须结合对应 Kimi 版本和实现核对。
12

主题讲解 · 03:37

MLA 如何用权重吸收绕过显式 QK 解压

学习目标

  • 能写出 MLA 的 Query latent 与 KV latent 投影。
  • 能用上投影列块表示不同 attention heads。
  • 能从 QrKrQ_rK_r^\top 推导吸收矩阵 MrM_r
  • 能检查 dcqdckvd_{cq}\ne d_{ckv} 时矩阵乘法仍然成立。
  • 能解释“无需显式解压”不等于没有投影计算。
  • 能明确权重吸收对 RoPE 内容/位置分支的适用边界。

前置与衔接

视频先画出 Query 与 KV 的显式压缩/解压路径。

设输入

XRN×dmodel.X\in\mathbb{R}^{N\times d_{model}}.

Query latent 与 KV latent 分别为

CQ=XWDQRN×dcq,C_Q=XW_{DQ} \in\mathbb{R}^{N\times d_{cq}},
CKV=XWDKVRN×dckv.C_{KV}=XW_{DKV} \in\mathbb{R}^{N\times d_{ckv}}.
图 1

输入 X 分别下投影为 CQ 与 CKV;图中同时画出显式上投影,便于先建立代数关系。

原视频 · 00:20 ↗

本课与视频一样,先省略 RoPE 的解耦位置分支。

推导描述的是 MLA 内容分数路径。

核心讲解

1. 显式解压的单头公式

对第 rr 个 head,令

WUQ,rRdcq×dh,W_{UQ,r}\in\mathbb{R}^{d_{cq}\times d_h},
WUK,rRdckv×dh.W_{UK,r}\in\mathbb{R}^{d_{ckv}\times d_h}.

显式解压得到

Qr=CQWUQ,r,Q_r=C_QW_{UQ,r},
Kr=CKVWUK,r.K_r=C_{KV}W_{UK,r}.

二者最后一维都为 dhd_h,才能计算点积。

2. 解压矩阵按列块产生不同 heads

完整 WUQW_{UQ} 可按输出列切成

WUQ=[WUQ,1WUQ,h].W_{UQ} =[W_{UQ,1}\mid\cdots\mid W_{UQ,h}].
图 2

把 WUQ 按输出列分块后,CQ 乘第 r 块即可得到第 r 个 Query 头。

原视频 · 01:20 ↗

CQC_Q 虽不带显式 head 轴,不同列块仍产生不同 QrQ_r

K 路径同理。

因此共享 latent 不代表所有 heads 完全相同。

3. 为什么它不是 MQA 的同一 KV 头

MQA 让所有 Query heads 直接读取同一组物理 K/V。

MLA 中每个 head 可使用自己的

WUQ,r,WUK,r.W_{UQ,r},W_{UK,r}.
图 3

不同 WUQ/WUK 列块生成不同 Q/K 头,因此共享压缩表示不等于 MQA 的相同 KV 头。

原视频 · 01:40 ↗

所以解压后的 Q/K heads 可彼此不同。

但它们都受共享低维 latent 的秩约束。

“保持多头逻辑”不等于与任意无约束 MHA 完全等价。

4. 把解压公式代入分数

内容分数为

QrKr.Q_rK_r^\top.

代入得到

(CQWUQ,r)(CKVWUK,r).(C_QW_{UQ,r}) (C_{KV}W_{UK,r})^\top.
图 4

将 Q_r=CQ·WUQ_r 与 K_r=CKV·WUK_r 代入 Q_rK_rᵀ,准备利用转置与结合律。

原视频 · 02:20 ↗

利用

(AB)=BA,(AB)^\top=B^\top A^\top,

可改写为

CQWUQ,rWUK,rCKV.C_QW_{UQ,r}W_{UK,r}^\top C_{KV}^\top.

5. 定义可预计算的吸收矩阵

Mr=WUQ,rWUK,r.M_r=W_{UQ,r}W_{UK,r}^\top.

MrRdcq×dckv.M_r\in\mathbb{R}^{d_{cq}\times d_{ckv}}.

模型参数固定后,MrM_r 可在加载或编译阶段预先融合。

推理时内容分数直接写为

Srcontent=CQMrCKVdh.S_r^{content} =\frac{C_QM_rC_{KV}^\top}{\sqrt{d_h}}.
图 5

固定参数 M_r=WUQ_r·WUK_rᵀ 可预先融合,使内容分数直接由 CQ·M_r·CKVᵀ 得到。

原视频 · 03:00 ↗

无需先物化高维 QrQ_rKrK_r 张量。

但仍要计算 latent、乘 MrM_r 与 latent 间的分数。

6. Shape 逐步检查

设 Query 长度为 NqN_q,KV 长度为 NkN_k

CQ:[Nq,dcq],C_Q:[N_q,d_{cq}],
Mr:[dcq,dckv],M_r:[d_{cq},d_{ckv}],
CKV:[dckv,Nk].C_{KV}^\top:[d_{ckv},N_k].

所以

[Nq,dcq][dcq,dckv][dckv,Nk]=[Nq,Nk].[N_q,d_{cq}] [d_{cq},d_{ckv}] [d_{ckv},N_k] =[N_q,N_k].

7. 两个 latent 宽度可以不同

显式解压后 Q/K 的 head width 必须相同,都是 dhd_h

但压缩空间只需由 MrM_r 连接:

dcqdckvd_{cq}\ne d_{ckv}

完全允许。

图 6

CQ 与 CKV 的压缩宽度可以不同;吸收矩阵 M_r 负责在两个 latent 空间之间映射。

原视频 · 03:20 ↗

MrM_r 是长方形时,矩阵乘法仍然自洽。

视频指出论文配置中 Query latent 可比 KV latent 更宽。

具体数值应以对应模型版本为准。

8. 为什么能改变执行顺序

原始显式路径先形成

Qr:[Nq,dh],Kr:[Nk,dh].Q_r:[N_q,d_h],\qquad K_r:[N_k,d_h].

吸收路径利用矩阵乘法结合律,先合并固定参数。

数学结果在精确算术下相同。

有限精度中,不同结合顺序可能产生微小舍入差异。

kernel 选择还会影响速度、显存和数值累积顺序。

9. RoPE 边界

完整 MLA 通常还包含解耦的 RoPE Query/Key 分量。

位置相关旋转不能一般性地完全吸收到一个固定 MrM_r 中。

因此完整分数常需合并:

  • 可权重吸收的内容分数;
  • 单独处理的位置分数。

视频板书明确省略 RoPE。

所以“无需显式解压 Q/K”应理解为该简化内容路径,而非所有分量都消失。

10. Value 路径不是本课证明对象

本推导只证明 QK 内容分数可绕过显式 Q/K 解压。

V 仍由 CKVC_{KV} 的上投影或等价的输出侧吸收参与计算。

不同 MLA kernel 可能采用不同结合方式。

不能从 QK 推导直接断言整个 attention 从不产生任何高维临时量。

跟练与练习

原视频定位

编者练习

给定 dcq=512,dckv=128,dh=64d_{cq}=512,d_{ckv}=128,d_h=64。 写出 WUQ,r,WUK,r,MrW_{UQ,r},W_{UK,r},M_r 的 shape,并验证 CQMrCKVC_QM_rC_{KV}^\top

查看参考答案

WUQ,rW_{UQ,r}[512,64][512,64]WUK,rW_{UK,r}[128,64][128,64]
Mr=WUQ,rWUK,rM_r=W_{UQ,r}W_{UK,r}^\top[512,128][512,128]
CQC_Q[Nq,512][N_q,512]CKVC_{KV}^\top[128,Nk][128,N_k],结果为 [Nq,Nk][N_q,N_k]

常见误区

  • 误区:共享 CQ/CKV 意味着所有 heads 完全相同。纠正:不同上投影列块生成不同逻辑 heads。
  • 误区:MLA 因此等价于任意无约束 MHA。纠正:共享 latent 带来低秩结构约束。
  • 误区:权重吸收后完全没有矩阵乘法。纠正:只是把固定上投影预融合,并在 latent 空间改变乘法顺序。
  • 误区:dcqd_{cq} 必须等于 dckvd_{ckv}。纠正:MrM_r 可以在两个不同宽度空间之间映射。
  • 误区:所有 MLA 分数都能被同一个固定矩阵吸收。纠正:RoPE 的位置分支需单独处理。
  • 误区:QK 可吸收就证明 V 路径也无需任何临时量。纠正:V 与输出侧还需独立分析。

本课小结

  • MLA 将 Query 与 KV 分别压缩为 CQ,CKVC_Q,C_{KV}
  • 不同 WUQ,r,WUK,rW_{UQ,r},W_{UK,r} 列块可从共享 latent 生成不同 heads。
  • 代入显式解压式后,QrKr=CQWUQ,rWUK,rCKVQ_rK_r^\top=C_QW_{UQ,r}W_{UK,r}^\top C_{KV}^\top
  • 预计算 Mr=WUQ,rWUK,rM_r=W_{UQ,r}W_{UK,r}^\top 后,可直接用 CQMrCKVC_QM_rC_{KV}^\top 得到内容分数。
  • dcqd_{cq}dckvd_{ckv} 可以不同,MrM_r 的 shape 负责连接两者。
  • 该结论限定于视频省略 RoPE 的内容路径;完整 MLA 仍需处理位置与 Value 分支。
13

主题讲解 · 03:14

因果 Mask 为什么能逐层阻断未来信息

学习目标

  • 能写出 causal mask 的逐元素定义。
  • 能证明 masked softmax 对未来位置给出零权重。
  • 能从 A×VA\times V 证明第 ii 行只依赖前缀 values。
  • 能用数学归纳法把单层因果性推广到任意层。
  • 能说明 RoPE、输出投影、归一化、FFN 与 residual 不破坏该结论的条件。
  • 能区分理想 -\infty 与低精度实现中的有限负数。

前置与衔接

对长度 NN 的 decoder self-attention,分数为

S=QKdk.S=\frac{QK^\top}{\sqrt{d_k}}.

定义 causal mask

Mij={0,ji,,j>i.M_{ij}= \begin{cases} 0,&j\le i,\\ -\infty,&j>i. \end{cases}

masked attention 为

A=softmaxrow(S+M).A=\operatorname{softmax}_{row}(S+M).
图 1

上三角 mask 阻断未来 token 的权重,使第 i 行输出只组合 V₁ 到 Vᵢ。

原视频 · 00:00 ↗

要证明“看不到未来”,不能只停在 mask 图案。

还要追踪 softmax、A×VA\times V 与后续所有子层的信息依赖。

核心讲解

1. 单行 masked softmax 的直接证明

对固定 Query 行 ii

Aij=exp(Sij+Mij)m=1Nexp(Sim+Mim).A_{ij} =\frac{\exp(S_{ij}+M_{ij})} {\sum_{m=1}^{N}\exp(S_{im}+M_{im})}.

j>ij>i,则

exp(Sij)=0.\exp(S_{ij}-\infty)=0.

因此

Aij=0,j>i.A_{ij}=0,\qquad j>i.

合法位置重新归一化:

j=1iAij=1.\sum_{j=1}^{i}A_{ij}=1.
图 2

理想 masked softmax 令 e^{-∞}=0,并只在 j≤i 的合法 key 上逐行归一化。

原视频 · 02:00 ↗

每行至少要有一个合法位置。

标准 causal mask 保留对角线,满足这一条件。

2. 分数缩放与 mask 是两个不同操作

分数按单头宽度缩放:

Sij=qikj/dk.S_{ij}=q_i^\top k_j/\sqrt{d_k}.

然后再施加可见性约束。

图 3

分数仍按单头宽度除以 sqrt(d_k),随后把 j>i 的位置屏蔽为理想 -∞。

原视频 · 01:40 ↗

1/dk1/\sqrt{d_k} 控制 logit 尺度。

causal mask 决定依赖边是否存在。

没有缩放仍可因果,只是数值与优化性质不同。

有缩放却没有 mask,也不会自动阻断未来。

3. A×V 把零权重变成依赖结论

ii 个输出为

oi=j=1NAijvj.o_i=\sum_{j=1}^{N}A_{ij}v_j.

由于 j>ij>iAij=0A_{ij}=0

oi=j=1iAijvj.o_i=\sum_{j=1}^{i}A_{ij}v_j.
图 4

A×V 的第 i 行是 V₁…Vᵢ 的加权和,不包含任何 j>i 的 value。

原视频 · 02:20 ↗

未来 vi+1,,vNv_{i+1},\ldots,v_N 在表达式中系数严格为零。

这才是“当前位置看不到未来”的直接代数含义。

4. RoPE 不增加跨 token 依赖边

RoPE 对位置 ii 的 Query 与位置 jj 的 Key 分别做确定性旋转。

分数可写成

(Riqi)(Rjkj)=qiRiRjkj.(R_iq_i)^\top(R_jk_j) =q_i^\top R_i^\top R_jk_j.
图 5

RoPE 改变允许位置上的 Q/K 内积,却不增加跨 token 的新输入边;因果 mask 仍决定可见性。

原视频 · 01:00 ↗

它会改变允许位置上的数值。

但位置 ii 的旋转只作用于 qiq_i,位置 jj 的旋转只作用于 kjk_j

RoPE 本身不把未来 token 特征注入过去 token 行。

随后 causal mask 仍把 j>ij>i 的连接置零。

5. 输出投影不跨序列行

多头输出拼接后经

yi=oiWO.y_i=o_iW_O.

WOW_O 混合特征通道与 heads。

它不会把其他 token 行 ojo_j 混入 oio_i

因此如果 oio_i 只依赖前缀,yiy_i 仍只依赖前缀。

6. FFN 不是线性的,但它逐 token

常见 FFN 为

FFN(xi)=W2ϕ(W1xi+b1)+b2.\operatorname{FFN}(x_i) =W_2\,\phi(W_1x_i+b_1)+b_2.

其中 ϕ\phi 是非线性激活。

视频口播把 FFN 概括为线性映射或逐位操作。

严格说 FFN 含非线性。

因果性真正依赖的是它对每个 token 行独立应用,而不是它必须线性。

7. Norm 与 residual 的条件

LayerNorm/RMSNorm 通常在单个 token 的特征维上归一化。

它们不沿序列轴聚合,所以不引入未来 token。

residual 也是同位置相加:

xinew=xi+fi(xi).x_i^{new}=x_i+f_i(x_{\le i}).

两项都只依赖前缀,和仍只依赖前缀。

如果人为引入跨 token pooling、非因果卷积或序列维 BatchNorm,该证明就需要重新检查。

8. 逐层归纳证明

归纳命题:第 \ell 层第 ii 行状态

xi()x_i^{(\ell)}

只依赖输入 token 1,,i1,\ldots,i

基例:embedding 行 xi(0)x_i^{(0)} 只来自 token ii 与其位置编码。

归纳假设:第 \ell 层所有 xj()x_j^{(\ell)} 只依赖各自前缀 1:j1{:}j

causal attention 的第 ii 行只读取 jij\le i 的 K/V。

这些 xj()x_j^{(\ell)} 依赖的最大输入位置不超过 ii

逐 token norm、FFN、输出投影与 residual 不增加序列依赖边。

所以 xi(+1)x_i^{(\ell+1)} 仍只依赖 1:i1{:}i

图 6

输出投影、残差、逐 token 归一化与 FFN 不跨序列行混合,因此该前缀依赖可逐层归纳保持。

原视频 · 03:00 ↗

由数学归纳法,任意深度都不混入未来信息。

9. 为什么训练能并行

训练时可以一次计算全部 NN 个位置。

并行执行不等于信息互相可见。

mask 在矩阵内部删除未来依赖边。

所以位置 ii 的 loss 可以与其他位置一起计算,却仍满足自回归因果约束。

10. 实际实现中的负大数

有限精度张量未必直接存 IEEE -\infty

常见实现使用 dtype 可表示的很大负数,或 fused masked softmax 在 kernel 内直接跳过无效项。

若用有限常数 C-C,理论上

eC>0.e^{-C}>0.

但在有限精度中常下溢为 0。

稳健实现应保证 masked 位置输出精确为 0 或在容差内为 0。

不能只凭常数“看起来很负”就忽略 dtype 与 kernel 行为。

11. Dropout 不应重新激活 masked 项

attention dropout 通常在 softmax 后对权重随机置零并重标定。

原本为 0 的 masked 项乘任何 dropout mask 仍为 0。

因此标准实现不破坏因果性。

若自定义算子在 mask 后又给全部位置加非零偏置,则必须重新审计。

跟练与练习

原视频定位

编者练习

长度为 4 时,写出理想 causal mask 的第二行,并说明 softmax 后第二行哪些权重一定为 0。

查看参考答案

第二行 mask 为
[0,0,,].[0,0,-\infty,-\infty].
加到原分数后逐行 softmax,第三、第四列的指数项为 0,所以 A2,3=A2,4=0A_{2,3}=A_{2,4}=0
第一、第二列在彼此之间重新归一化,总和为 1。

常见误区

  • 误区:分数除以 dk\sqrt{d_k} 就能保证因果。纠正:因果性来自 mask,缩放只控制数值尺度。
  • 误区:FFN 必须线性才不泄漏未来。纠正:它可以非线性,关键是逐 token 应用。
  • 误区:RoPE 会跨位置旋转,所以会混合 token。纠正:它改变位置相关内积,不把一个位置的特征写入另一个位置。
  • 误区:只证明一层 softmax 为零就结束。纠正:还要验证 A×VA\times V 与后续所有操作不新增跨行依赖。
  • 误区:任何有限大负数都等价于 -\infty。纠正:需结合 dtype、下溢和 fused kernel 验证。
  • 误区:并行训练意味着位置能看到未来。纠正:计算并行与依赖图的因果可见性是两回事。

本课小结

  • causal mask 对 j>ij>i 加理想 -\infty,masked softmax 令这些权重为 0。
  • ii 行输出因此只包含 V1,,ViV_1,\ldots,V_i 的加权组合。
  • RoPE 与 1/dk1/\sqrt{d_k} 改变允许边的数值,不改变 mask 定义的依赖图。
  • 输出投影、逐 token norm、FFN 与 residual 都不跨序列行混合。
  • 由逐层归纳,任意深度的第 ii 行都只依赖输入前缀 1:i1{:}i
  • 实际实现应审计有限负数、低精度与 fused masked softmax 是否保持屏蔽项为零。
14

主题讲解 · 02:50

Query 与 Key 的 Token 数如何决定注意力矩阵形状

学习目标

  • 能用 Nq,NkN_q,N_k 一步判断任意注意力分数矩阵 shape。
  • 能说明全序列 self-attention 为什么通常是正方形。
  • 能说明 cross-attention 为什么通常是非正方形长方形。
  • 能识别 cross-attention 偶然为正方形的情况。
  • 能识别增量 self-attention 为 1×L1\times L 的反例。
  • 能用 A×VA\times V 的收缩轴验证输出长度。

前置与衔接

对单个 attention head,设

QRNq×dk,Q\in\mathbb{R}^{N_q\times d_k},
KRNk×dk.K\in\mathbb{R}^{N_k\times d_k}.

分数为

S=QKdk.S=\frac{QK^\top}{\sqrt{d_k}}.

矩阵乘法直接给出

SRNq×Nk.S\in\mathbb{R}^{N_q\times N_k}.

这是一条比“self 是方形、cross 是长方形”更一般的规则。

核心讲解

1. 每行与每列分别代表什么

分数元素

Sij=qikjdkS_{ij}=\frac{q_i^\top k_j}{\sqrt{d_k}}

表示第 ii 个 Query 对第 jj 个 Key 的匹配分数。

所以:

  • 行数等于 Query token 数 NqN_q
  • 列数等于 Key token 数 NkN_k

注意力矩阵是否正方形,只由 Nq=NkN_q=N_k 是否成立决定。

2. 全序列 self-attention 通常是正方形

全序列 self-attention 从同一批 NN 个 token 同时生成 Q 与 K。

Nq=Nk=N.N_q=N_k=N.

于是

SRN×N.S\in\mathbb{R}^{N\times N}.
图 1

全序列 self-attention 的 Q/K 来自同一批 N 个 token,因此分数矩阵通常为 N×N。

原视频 · 00:20 ↗

“同源”常常导致长度相同。

但真正决定 shape 的仍是参与本次计算的 Query/Key 长度。

3. 三 token 的具体例子

视频令 decoder 输入包含三个 token。

拆头后每个 head 有

QR3×dk,Q\in\mathbb{R}^{3\times d_k},
KRdk×3.K^\top\in\mathbb{R}^{d_k\times3}.

所以

QKR3×3.QK^\top\in\mathbb{R}^{3\times3}.
图 2

三个 decoder token 同时产生三行 Q 与三列 Kᵀ,形成 3×3 masked self-attention。

原视频 · 01:20 ↗

decoder 还会在这个 3×33\times3 矩阵上施加 causal mask。

mask 改变允许的连接,不改变矩阵 shape。

4. Cross-attention 的两侧长度独立

encoder-decoder cross-attention 中:

Q=HdecWQ,Q=H_{dec}W_Q,
K=HencWK,K=H_{enc}W_K,
V=HencWV.V=H_{enc}W_V.
图 3

交叉注意力的 Q 来自 decoder 的三行状态,K/V 来自 encoder 的五行 memory。

原视频 · 01:40 ↗

decoder 长度 NdecN_{dec} 与 encoder 长度 NencN_{enc} 可独立变化。

所以一般有

ScrossRNdec×Nenc.S^{cross} \in\mathbb{R}^{N_{dec}\times N_{enc}}.

5. 三个 Query 对五个 Key

视频例子取

Nq=3,Nk=5.N_q=3,\qquad N_k=5.

因此

[3,dk][dk,5]=[3,5].[3,d_k][d_k,5]=[3,5].
图 4

三行 Q 乘五列 Kᵀ,得到 shape 为 3×5 的交叉注意力矩阵。

原视频 · 02:00 ↗

每一行表示一个文本 Query 如何在五个 encoder token 上分配权重。

每一列对应一个 encoder Key 被三个文本 Query 读取的分数。

6. Cross-attention 也可以是正方形

如果恰好

Ndec=Nenc,N_{dec}=N_{enc},

cross-attention 的 shape 也是 N×NN\times N

图 5

若 Query 与 encoder memory 的 token 数恰好相等,交叉注意力也会得到正方形矩阵。

原视频 · 00:40 ↗

因此不能仅凭矩阵外形判断 attention 类型。

还必须看 Q 与 K/V 的来源。

日常说“cross-attention 是长方形”通常是指两侧长度经常不同。

若把“长方形”严格用作非正方形,这句话存在例外。

7. Self-attention 也可以是长方形

自回归增量 decode 时,当前步 Query 长度常为 1。

KV Cache 中已有 LL 个 Key/Value 位置。

于是 self-attention 分数为

[B,h,1,dk]×[B,h,dk,L][B,h,1,L].[B,h,1,d_k] \times [B,h,d_k,L] \rightarrow [B,h,1,L].

Q 与 K 仍属于同一个序列和同一个 self-attention 模块。

但本次计算的长度不同,所以矩阵是 1×L1\times L

这再次说明 shape 规则应从 Nq,NkN_q,N_k 出发。

8. A×V 为什么输出跟随 Query 长度

softmax 后

ARNq×Nk.A\in\mathbb{R}^{N_q\times N_k}.

Value 为

VRNk×dv.V\in\mathbb{R}^{N_k\times d_v}.

所以

O=AVRNq×dv.O=AV \in\mathbb{R}^{N_q\times d_v}.
图 6

3×5 的注意力权重与五行 V 相乘,收缩 key/value 长度,输出恢复为三个 Query 行。

原视频 · 02:20 ↗

NkN_k 是被收缩的求和轴。

输出为每一个 Query 生成一个 dvd_v 维向量。

9. 加入 batch 与 head 轴

一般多头 shape 为

QRB×hq×Nq×dk,Q\in\mathbb{R}^{B\times h_q\times N_q\times d_k},
KRB×hkv×Nk×dk.K\in\mathbb{R}^{B\times h_{kv}\times N_k\times d_k}.

在 MHA 中 hq=hkvh_q=h_{kv}

在 GQA/MQA 中需先建立 Query 头到 KV 头的映射。

无论头如何映射,每个 Query 头的分数二维仍是

Nq×Nk.N_q\times N_k.

10. Mask 不改变矩阵外形

self-attention 的 causal mask、padding mask,或 cross-attention 的 source mask,都会广播到分数 shape。

它们把部分位置设为不可见。

但不会删除矩阵行列。

因此“上三角是零”与“矩阵仍是方形”可以同时成立。

跟练与练习

原视频定位

编者练习

当前 decode 步有两个 Query,KV Cache 含 512 个历史位置,每头 dk=64d_k=64。 写出单头 Q、Kᵀ、分数与输出的 shape。

查看参考答案

Q 为 [2,64][2,64],Kᵀ 为 [64,512][64,512],分数为 [2,512][2,512]
若 V 为 [512,dv][512,d_v],输出为 [2,dv][2,d_v]
这是 self-attention,但本次 Q/K 长度不同,因此分数不是正方形。

常见误区

  • 误区:所有 self-attention 都必为正方形。纠正:增量 decode 可为 1×L1\times L
  • 误区:所有 cross-attention 都必为非正方形。纠正:两侧长度相等时也会是方形。
  • 误区:从外形就能判断 self 或 cross。纠正:类型取决于 Q 与 K/V 的来源。
  • 误区:长方形 A 与 V 无法相乘。纠正:[Nq,Nk][Nk,dv][N_q,N_k][N_k,d_v] 的中间轴完全匹配。
  • 误区:mask 会把矩阵裁成三角形 shape。纠正:shape 不变,只是部分位置被屏蔽。
  • 误区:输出长度跟随 Key。纠正:Key 长度被收缩,输出行数跟随 Query。

本课小结

  • 任意单头注意力分数的 shape 都是 Nq×NkN_q\times N_k
  • 全序列 self-attention 常有 Nq=NkN_q=N_k,所以通常是正方形。
  • cross-attention 两侧长度独立,所以通常是非正方形长方形。
  • cross-attention 可因等长而为方形,增量 self-attention 也可因 Q/K 不等长而为长方形。
  • A×VA\times V 收缩 NkN_k,输出保留 NqN_q
  • 判断 attention 类型要看信息来源,判断矩阵形状要看实际 Query/Key token 数。
15

单元综合

从注意力计算骨架到高效结构变体

单元能力目标

完成本单元后,应能把 attention 看成一条可以逐轴检查、逐步变形的计算链,而不是只记住一个公式。

具体需要做到:

  • 从 Query/Key/Value 的来源与 shape 判断 self-attention、cross-attention 和增量 attention;
  • 写出拆头、逐头打分、softmax、Value 汇聚、拼头与输出投影的完整张量路径;
  • 解释缩放、温度、causal mask 与 attention sink 分别改变哪一部分数学语义;
  • hqh_qhkvh_{kv} 和潜变量宽度统一比较 MHA、GQA、MQA 与 MLA;
  • 从缓存张量而不是模型名称估算 decode 的容量与带宽开销;
  • 推导标准 attention 对 Q,K,VQ,K,V 的完整反向传播;
  • 说明权重吸收为何能改变 MLA 的执行顺序,以及 RoPE 为何构成版本边界;
  • 区分沿 token 轴做注意力与沿历史分支轴做 Attention Residual;
  • 对任何新 attention 变体先建立 shape、归一化轴、可见性、缓存落点和参数化五项台账。

概念连接

1. 一条统一的注意力计算链

对单个 head,设

QRNq×dk,KRNk×dk,VRNk×dv.Q\in\mathbb{R}^{N_q\times d_k},\qquad K\in\mathbb{R}^{N_k\times d_k},\qquad V\in\mathbb{R}^{N_k\times d_v}.

标准计算可写为

S=QKdk+M,S=\frac{QK^\top}{\sqrt{d_k}}+M,
P=softmaxNk(S),P=\operatorname{softmax}_{N_k}(S),
O=PV.O=PV.

shape 沿这条链变化为

[Nq,dk][dk,Nk][Nq,Nk],[N_q,d_k][d_k,N_k] \rightarrow[N_q,N_k],
[Nq,Nk][Nk,dv][Nq,dv].[N_q,N_k][N_k,d_v] \rightarrow[N_q,d_v].

这两次收缩给出三个最稳定的判断规则:

  • 分数矩阵行数跟随 Query;
  • 分数矩阵列数跟随 Key;
  • 输出行数仍跟随 Query,因为 Key/Value 长度在 PVPV 中被收缩。

因此,矩阵是方形还是长方形只由本次计算的 NqN_qNkN_k 决定。

全序列 self-attention 常有 Nq=Nk=NN_q=N_k=N,所以常见 N×NN\times N;cross-attention 两侧长度独立,所以常见 Nt×NsN_t\times N_s;增量 self-attention 则可直接是 1×L1\times L

不能只凭矩阵外形判断 self 或 cross。类型取决于 Q 与 K/V 的来源,shape 取决于实际参与本次计算的 token 数。

2. 多头机制先独立计算,再重新混合

标准 MHA 先把投影特征组织为

QRB×h×Nq×dk,Q\in\mathbb{R}^{B\times h\times N_q\times d_k},
KRB×h×Nk×dk,K\in\mathbb{R}^{B\times h\times N_k\times d_k},
VRB×h×Nk×dv.V\in\mathbb{R}^{B\times h\times N_k\times d_v}.

每个 head 独立计算 QrKrQ_rK_r^\topPrVrP_rV_r

softmax 沿 Key 位置轴 NkN_k 做,而不是沿 head 轴做。

各头输出随后交换轴并拼接:

OcatRB×Nq×hdv.O_{cat}\in\mathbb{R}^{B\times N_q\times hd_v}.

输出投影

Y=OcatWOY=O_{cat}W_O

会重新混合不同 heads 的特征。

所以“各头独立”只适用于逐头 attention 阶段,不适用于整个注意力层的最终输出。

3. 缩放维度来自点积收缩轴

单个分数是

s=qk=c=1dkqckc.s=q^\top k=\sum_{c=1}^{d_k}q_ck_c.

在各维独立、零均值、单位方差的教学假设下,

E[s]=0,Var(s)=dk.\mathbb{E}[s]=0, \qquad \operatorname{Var}(s)=d_k.

于是

Var(sdk)=1.\operatorname{Var}\left(\frac{s}{\sqrt{d_k}}\right)=1.

缩放使用 dkd_k,因为每个 head 的点积只收缩 dkd_k 个坐标。

在常见 dmodel=hdkd_{model}=h d_k 配置下,若误除以 dmodel\sqrt{d_{model}},理想方差会被压到 1/h1/h

这一结论依赖简化统计假设。训练后的 Q/K 不必严格独立同分布;固定缩放校正的是随点积维度增长的典型尺度,不保证注意力一定均匀。

4. 缩放、温度、稳定 softmax 是三件相邻但不同的事

softmax 温度写作

pi(T)=ezi/Tjezj/T.p_i(T)=\frac{e^{z_i/T}}{\sum_j e^{z_j/T}}.

从代数上看,1/dk1/\sqrt{d_k} 是对未缩放点积使用固定温度 T=dkT=\sqrt{d_k}

但设计上要区分:

  • 1/dk1/\sqrt{d_k} 用来补偿点积尺度随 head width 增长;
  • 额外温度用于主动改变真实 token 之间的概率比例;
  • softmax 减去行最大值用于避免指数溢出,不负责点积方差校准。

温度改变任意两项的概率比:

pi(T)pj(T)=exp(zizjT).\frac{p_i(T)}{p_j(T)} =\exp\left(\frac{z_i-z_j}{T}\right).

高温通常让分布更平,低温通常让分布更尖;正温度不改变 logit 排序。

5. Attention Sink 改变的是总上下文质量

对真实 logits z1,,znz_1,\ldots,z_n,令

Z=j=1nezj.Z=\sum_{j=1}^{n}e^{z_j}.

加入 sink logit ss 后,真实 token 权重为

ai=eziZ+es.a_i=\frac{e^{z_i}}{Z+e^s}.

若普通权重为 πi=ezi/Z\pi_i=e^{z_i}/Z,则

ai=gπi,g=ZZ+es.a_i=g\pi_i, \qquad g=\frac{Z}{Z+e^s}.

所以单独增加 sink 不改变真实 token 之间的相对比例,只把它们共同乘以门控 gg

若 sink value 为零,输出为

osink=giπivi.o_{sink}=g\sum_i\pi_i v_i.

这与温度的区别可以概括为:温度在真实 token 内部重新分配固定总量,sink 把一部分总量导向额外槽位。

若 sink value 非零,输出还会加入该额外向量;若温度也作用于 sink logit,两种机制会发生耦合,必须先写清参数化再比较。

6. Causal mask 删除依赖边

理想 causal mask 为

Mij={0,ji,,j>i.M_{ij}= \begin{cases} 0,&j\le i,\\ -\infty,&j>i. \end{cases}

于是 masked softmax 满足

Pij=0,j>i.P_{ij}=0,\qquad j>i.

ii 个输出化为

oi=j=1iPijvj.o_i=\sum_{j=1}^{i}P_{ij}v_j.

这才是“看不到未来”的代数含义。

缩放只改变允许边上的分数尺度,RoPE 只改变位置相关匹配数值,都不会代替 causal mask 删除未来边。

输出投影、逐 token RMSNorm/LayerNorm、逐 token FFN 与同位置 residual 不跨序列行,因此可用数学归纳法把前缀依赖推广到任意深度。

实际 kernel 可能用有限大负数或直接跳过屏蔽项。工程验收应检查 masked 权重在目标 dtype 下确为零或在规定容差内为零。

7. Cross-attention 只改变 Q 与 K/V 的来源

经典 encoder-decoder cross-attention 使用

Q=HdecWQcross,Q=H_{dec}W_Q^{cross},
K=HencWKcross,V=HencWVcross.K=H_{enc}W_K^{cross}, \qquad V=H_{enc}W_V^{cross}.

若 decoder 长度为 NtN_t、encoder memory 长度为 NsN_s,则

PRB×h×Nt×Ns,P\in\mathbb{R}^{B\times h\times N_t\times N_s},
ORB×h×Nt×dv.O\in\mathbb{R}^{B\times h\times N_t\times d_v}.

它通常没有 decoder self-attention 的上三角 causal mask,但仍可有 encoder padding mask、有效 patch mask 或任务定义的稀疏 mask。

每个 decoder 层可读取同一份 encoder memory,却通常使用各自的 cross-attention 投影参数。

8. MHA、GQA 与 MQA 是 KV 头数上的离散轴

固定 Query 头数 hqh_q,用 KV 头数 hkvh_{kv} 统一描述:

结构hkvh_{kv}每组服务的 Query 头数
MHAhqh_q11
GQA1<hkv<hq1<h_{kv}<h_qg=hq/hkvg=h_q/h_{kv}
MQA11hqh_q

若采用等大小连续分组,第 rr 个 Query head 可映射到

m(r)=rg.m(r)=\left\lfloor\frac{r}{g}\right\rfloor.

共享 K/V 不代表各头输出相同,因为 QrQ_r 仍可不同:

Pr=softmax(QrKm(r)dk+Mr).P_r=\operatorname{softmax}\left( \frac{Q_rK_{m(r)}^\top}{\sqrt{d_k}}+M_r \right).

只要 Query 不同,分数、概率与输出通常也不同。

逻辑广播只描述多个 Query heads 读取同一份 K/V;高效实现不应为了套普通 MHA 接口而物化重复缓存。

9. KV Cache 开销首先由 h_kv 决定

忽略元数据时,一层、一个样本、长度 LL 的缓存主体约为

NKV=2Lhkvdh.N_{KV}=2Lh_{kv}d_h.

所以在其余条件不变时,GQA 相对 MHA 的主体比例为

hkvhq,\frac{h_{kv}}{h_q},

MQA 则约为

1hq.\frac{1}{h_q}.

这个比例只描述 K/V 张量主体,不等于端到端显存或端到端延迟比例。

页表、对齐、量化尺度、模型权重、临时工作区和 kernel 利用率都可能改变最终收益。

decode 常受长上下文 K/V 读取带宽限制,因此减少物理 KV heads 往往比减少同阶算术更直接影响吞吐。

10. MLA 改为缓存低维潜变量

简化 MLA 先计算

CKV=XWDKVRN×dckv.C_{KV}=XW_{DKV} \in\mathbb{R}^{N\times d_{ckv}}.

decode 跨步缓存 CKVC_{KV},而不是完整 K/V heads。

显式教学路径可写为

Kr=CKVWUK,r,Vr=CKVWUV,r.K_r=C_{KV}W_{UK,r}, \qquad V_r=C_{KV}W_{UV,r}.

不同上投影列块仍可生成不同逻辑 heads,所以共享 latent 不等于 MQA 的同一物理 K/V head。

但合成投影

WK=WDKVWUKW_K=W_{DKV}W_{UK}

的秩至多为 latent 宽度,因此输出 shape 与 MHA 相同不代表可表达任意无约束 MHA 投影。

忽略位置分支时,传统 K/V 每 token 的元素量约为 2hkvdh2h_{kv}d_h,latent cache 主体约为 dckvd_{ckv};完整实现还可能保存解耦 RoPE 分量。

11. 权重吸收把固定上投影移出逐 token 路径

若 Query 也先压缩为 CQC_Q,第 rr 个 head 的内容分数为

QrKr=(CQWUQ,r)(CKVWUK,r).Q_rK_r^\top =(C_QW_{UQ,r})(C_{KV}W_{UK,r})^\top.

利用转置规则与结合律:

QrKr=CQWUQ,rWUK,rCKV.Q_rK_r^\top =C_QW_{UQ,r}W_{UK,r}^\top C_{KV}^\top.

定义固定吸收矩阵

Mr=WUQ,rWUK,r,M_r=W_{UQ,r}W_{UK,r}^\top,

即可直接计算

Srcontent=CQMrCKVdh.S_r^{content} =\frac{C_QM_rC_{KV}^\top}{\sqrt{d_h}}.

即使 dcqdckvd_{cq}\ne d_{ckv},只要

MrRdcq×dckv,M_r\in\mathbb{R}^{d_{cq}\times d_{ckv}},

shape 仍然成立。

“无需显式解压 Q/K”表示改变矩阵乘法顺序,不表示计算或投影消失。

位置相关 RoPE 旋转通常不能整体吸收到固定矩阵中;Value 与输出侧也要单独分析。该推导只覆盖视频简化的内容分数路径。

12. 反向传播沿两次矩阵乘法分叉

前向记为

S=QK/dk+M,S=QK^\top/\sqrt{d_k}+M,
P=softmax(S),O=PV.P=\operatorname{softmax}(S), \qquad O=PV.

设上游梯度为 G=L/OG=\partial\mathcal{L}/\partial O

O=PVO=PV

dP=GV,dV=PG.dP=GV^\top, \qquad dV=P^\top G.

softmax 反向必须先把 dPdP 变成 dSdS。对每一行:

dS=P(dProwsum(dPP)).dS =P\odot\left( dP-\operatorname{rowsum}(dP\odot P) \right).

再由分数路径得到

dQ=dSKdk,dK=dSQdk.dQ=\frac{dS K}{\sqrt{d_k}}, \qquad dK=\frac{dS^\top Q}{\sqrt{d_k}}.

五矩阵图适合记依赖骨架,但不能把 softmax 前 logits SS 与 softmax 后概率 PP 混写,也不能漏掉 dVdV、scale 与 mask 边界。

FlashAttention 可通过分块、在线 softmax 与反向重计算不保存完整 S/PS/P,但目标梯度仍与这条数学链等价。

13. Attention Residual 把注意力轴转向历史分支

普通 token attention 沿 Key token 轴归一化。

视频所示 Attention Residual 对固定 token 收集 mm 条 clean branch 输出:

VhistRN×m×d.V_{hist}\in\mathbb{R}^{N\times m\times d}.

历史向量可作为 value,其 RMSNorm 版本作为 key;层级 residual query 与这些 keys 匹配:

αi,j=softmaxj((qres)ki,j),\alpha_{i,j} =\operatorname{softmax}_{j} \left((q_{res})^\top k_{i,j}\right),
yi=j=1mαi,jvi,j.y_i=\sum_{j=1}^{m}\alpha_{i,j}v_{i,j}.

这里 softmax 轴是同一 token 的历史 branch,而不是序列 token。

层级 query 可以共享,但历史 keys 随 token 变化,因此权重仍可逐 token 动态变化。

clean branch 提供直接历史路径,是减轻嵌套残差稀释的设计动机,不是对任意深度的信息保真定理。

对比与决策

1. 先判断要改的是哪一层语义

问题直接机制保持不变的关键部分
点积随 head width 变大除以 dk\sqrt{d_k}token 排序不由固定正缩放改变
想改变真实 token 间尖锐度softmax 温度真实 token 总概率仍为 1
想允许某个 head 少读上下文sink logit真实 token 间比例可保持不变
想阻断未来 tokencausal mask分数矩阵 shape 不变
想读取另一序列cross-attention注意力仍是 QKPVQK^\top\rightarrow PV
想减少 KV headsGQA/MQAQuery 多头与逐头权重仍可保留
想缓存更低维状态MLA解压后逻辑 heads 仍可不同
想读取历史层分支Attention Residual每个 token 的输出 shape 仍为 dd

2. 选择 MHA、GQA、MQA 或 MLA 时看四项

第一,看质量与容量需求。

  • MHA 给 K/V 侧最高的逐头独立度;
  • GQA 在若干 Query heads 内共享 K/V;
  • MQA 把 K/V 侧压到一组;
  • MLA 用共享低维 latent 生成或吸收多个逻辑 heads。

第二,看 decode 缓存主体。

  • MHA/GQA/MQA 主要按 hkvh_{kv} 计;
  • MLA 主要按 latent 宽度与额外位置分量计。

第三,看 kernel 是否真正利用共享。

  • 逻辑广播若被物化为重复 K/V,会抵消 GQA/MQA 的优势;
  • MLA 若总是显式解压完整 K/V,也可能放弃部分权重吸收收益。

第四,看版本边界。

  • RoPE 分支、Q/K latent 宽度、Value 吸收、量化与分页缓存都会改变真实实现;
  • 不能把简化板书直接当成某个生产模型的完整规范。

3. 审查任意 attention 变体的五问

  1. Q、K、V 各来自哪个状态序列?
  2. 每次点积的收缩维、Query 长度和 Key 长度是什么?
  3. softmax 沿哪个轴,mask 删除哪些依赖边?
  4. decode 跨步真正缓存哪个张量,是否物化广播或解压结果?
  5. 哪些变换是固定线性参数,可以用结合律预融合;哪些含位置、非线性或动态状态,不能直接吸收?

只要这五项写清,名称相似但机制不同的结构就不容易混淆。

综合训练

编者练习

设计一个 decoder attention 层,配置如下: 回答:

  1. batch size 为 BB
  2. 当前增量 Query 长度为 1;
  3. KV Cache 已含 LL 个可见位置;
  4. Query 头数 hq=32h_q=32
  5. KV 头数 hkv=8h_{kv}=8
  6. 单头宽度 dk=dv=128d_k=d_v=128
  7. 每个 head 还加入一个 value 为零的 sink logit;
  8. 使用标准 causal mask 与缩放点积;
  9. 上游输出梯度 shape 与 head 输出相同。
  10. 这是 MHA、GQA 还是 MQA?每组有多少 Query heads?
  11. 拆头后 Q、物理 K/V、分数、概率与 head 输出的 shape 是什么?
  12. 不计 dtype 与元数据,每层 KV Cache 主体有多少元素?相对同宽 MHA 的比例是多少?
  13. sink 对真实 token 概率比与真实 token 总质量分别有什么影响?
  14. 写出从 dOdOdV,dP,dS,dQ,dKdV,dP,dS,dQ,dK 的反向顺序。
查看参考答案

• 因为 1<hkv<hq1<h_{kv}<h_q,这是 GQA。group size 为 g=hq/hkv=32/8=4g=h_q/h_{kv}=32/8=4,每组四个 Query heads 共享一个物理 KV head。
• 当前步 Query 为 Q:[B,32,1,128]Q:[B,32,1,128];物理 K/V 为 K,V:[B,8,L,128]K,V:[B,8,L,128]。建立 Query head 到 KV head 的映射后,逻辑分数与概率为 S,P:[B,32,1,L]S,P:[B,32,1,L],每个 Query head 的输出整体为 O:[B,32,1,128]O:[B,32,1,128]。sink 是归一化中的额外槽位,不改变物理 K/V 长度 LL
• 每层缓存主体元素量为 2BLhkvdh=2BL×8×128=2048BL2BLh_{kv}d_h=2BL\times8\times128=2048BL。同宽 MHA 使用 hkv=hq=32h_{kv}=h_q=32,所以比例为 8/32=1/48/32=1/4
• 设真实 logits 指数和为 ZZ、sink logit 为 ss。真实 token 权重共同乘 gsink=Z/(Z+es)g_{sink}=Z/(Z+e^s)。它们之间的概率比保持不变,但总质量从 1 降为 gsinkg_{sink};sink value 为零时,真实上下文 value 混合也整体乘该门控。
• 由 O=PVO=PVdP=dOVdP=dO\,V^\topdV=PdOdV=P^\top dO。再逐行通过包含 sink 槽位的 softmax Jacobian 得到 dSdS;若 sink logit 可训练,还要得到其梯度。随后 dQ=dSK/dkdQ=dS K/\sqrt{d_k}dK=dSQ/dkdK=dS^\top Q/\sqrt{d_k}。GQA 共享物理 K/V,因此同一 KV group 内多个 Query heads 的 dK,dVdK,dV 要累加回对应物理 head。

编者练习

某简化 MLA 内容路径有 CQ:[2,512]C_Q:[2,512]CKV:[1024,128]C_{KV}:[1024,128],单个 head 的上投影为 WUQ:[512,64]W_{UQ}:[512,64]WUK:[128,64]W_{UK}:[128,64]。 推导吸收矩阵与最终内容分数的 shape,并指出为什么该推导不能自动覆盖 RoPE 和 Value 路径。

查看参考答案

吸收矩阵为 M=WUQWUK:[512,128]M=W_{UQ}W_{UK}^\top:[512,128]
于是 CQMCKVC_QMC_{KV}^\top 的 shape 链为 [2,512][512,128][128,1024]=[2,1024][2,512][512,128][128,1024]=[2,1024]
该等式依赖固定线性上投影与矩阵乘法结合律。RoPE 含位置相关旋转,通常不能全部并入一个固定 MM;Value 的上投影与 PVPV、输出投影如何结合需要另一条独立推导,不能由 QK 内容分数的等式直接推出。

进入下一单元前

  • 已能从 Nq,Nk,dk,dv,hq,hkvN_q,N_k,d_k,d_v,h_q,h_{kv} 写出任意注意力主路径的 shape。
  • 已能把缩放、温度、sink 与 mask 分别定位到“尺度、比例、总质量、可见性”四种不同作用。
  • 已能从物理缓存张量解释 MHA、GQA、MQA 与 MLA 的效率差异。
  • 已能完整写出 dV,dP,dS,dQ,dKdV,dP,dS,dQ,dK,不再把 logits 与概率混用。
  • 已能识别权重吸收的线性代数条件,并主动标记 RoPE 与 Value 路径边界。
  • 仍不稳定时,先回看多头 shape、单头缩放、矩阵形状与五矩阵反向骨架,再进入位置编码。
  • 下一单元将把注意力分数中的位置因素单独展开:从二维旋转、复数欧拉公式和旋转矩阵复合,推导 RoPE 为什么只保留相对位移,以及不同代码布局如何实现同一旋转语义。