2017 年的 Attention Is All You Need 提出了 Transformer。论文以机器翻译中的 sequence transduction 为主要任务,用 attention-based 的 encoder-decoder 结构替代当时常见的循环网络和卷积网络。它的核心主张是:序列位置之间的依赖可以主要通过 attention 建模,而不必依赖沿时间步递推的 RNN。这篇文章要解释的是 Transformer 里 Attention 算子本身如何工作,以及它为什么能成为后续大模型结构的核心部件

上一篇 Attention 前置:从序列瓶颈到内积、Softmax,再到 GPU 内存墙 介绍了必要的背景:序列不适合被压缩成一个固定状态,token 需要连续向量表示,内积可以把两个向量的匹配程度变成分数,Softmax 可以把一组分数转换成非负且总和为 1 的权重,GPU 上的实现还会受 HBM 读写限制。本文从这些前置概念出发,解释公式为什么写成 softmax(QKT/dk)Vsoftmax(QK^T / \sqrt{d_k})V本文的目标不是背下 Attention 公式,而是理解公式里每个符号背后的含义

这篇文章沿着两条线路展开。前半部分严格对应原论文的核心结构:Scaled Dot-Product Attention、mask、Multi-Head Attention、位置编码、Transformer block;后半部分讨论现代 decoder-only LLM 中围绕同一 Attention 算子形成的推理实践,包括 KV Cache、MQA/GQA、FlashAttention 和 PagedAttention。原论文给出的是 Transformer 和 Attention 的基础形式,现代推理优化解决的是同一算子进入长上下文和在线服务场景后的执行成本

1. 从任务到一次信息读取

Attention Is All You Need 这篇论文的实验主线是机器翻译,但 attention function 本身不是翻译专用模块,而是一种更通用的序列建模机制。这个机制在不同任务里对应不同的读取动作:例如机器翻译中,decoder 当前位置要从 source sentence 里读取相关词;自回归语言模型预测下一个 token 时,当前位置要从已有上下文里读取相关 token。无论任务是翻译还是续写,Attention 解决的都是“当前位置应该从哪些位置读取信息、各读取多少”的问题

先看一个具体句子:

小明 把 苹果 放进 书包,因为 它 很重。

当模型处理到“它”这个位置时,单靠“它”本身无法判断指代对象。当前位置需要回头看前面的“小明”“苹果”“书包”等 token,并结合“很重”这个描述判断哪些信息更相关。这个过程可以理解成一次软检索:当前位置提出问题,所有候选位置都参与匹配,但每个位置贡献的比例不同

这里的 query、key 和 value 不是人工标注的标签,而是当前层输入表示经过三组可训练矩阵投影得到的中间向量。训练时,反向传播会同时更新 embedding、WQW_QWKW_KWVW_V 等参数,使这些向量逐渐服务于最终的预测目标。QKV 的语义来自训练结果,而不是模型结构里预先写死的规则

在这个例子里,“它”这个位置发出的 query 可以直观理解为:“我需要找到一个前面出现过、可以被代词指代、并且和重量描述有关的对象。”前文每个 token 的 key 则像是它暴露出来的匹配标签:小明更像人名和主体,苹果是物体和食物,书包是物体和容器。Query 和 Key 负责计算匹配程度,也就是决定“看谁、看多少”

被匹配的位置真正贡献给当前 token 的不是 key,而是 value。苹果的 value 携带关于“苹果”的语义内容,书包的 value 携带关于“书包”的语义内容,小明的 value 携带关于“小明”的语义内容。Key 用来寻址,Value 用来传递内容,这就是 QKV 拆分最重要的直觉

模型可能给这些位置分配如下权重:

小明:0.05
苹果:0.65
书包:0.25
其他:0.05

那么“它”这个位置的新表示就不是直接选择某一个词,而是对所有 value 做加权和:

o=0.05v小明+0.65v苹果+0.25v书包+0.05v其他\begin{aligned} o_{\text{它}} &= 0.05 \cdot v_{\text{小明}} + 0.65 \cdot v_{\text{苹果}} + 0.25 \cdot v_{\text{书包}} + 0.05 \cdot v_{\text{其他}} \end{aligned}

这个输出仍然是一个向量,但它已经吸收了上下文中更相关的位置的信息。写成单个 query 对一组 key-value pairs 的形式,就是先用 query 和每个 key 打分,再用 Softmax 把分数变成权重,最后对 values 求加权和。Scaled Dot-Product Attention 正是把这次信息读取矩阵化:用 QKTQK^T 同时打出所有位置对的分数,用 dk\sqrt{d_k} 控制尺度,用 Softmax 得到权重,再乘以 VV 聚合内容

2. QKV 角色拆分

上一篇已经用 XXTXX^T 说明:只要把一段序列写成矩阵 XRN×DX \in \mathbb{R}^{N \times D},就可以通过内积一次性计算所有 token 对之间的相关性。但直接用 XXTXX^T 做 attention 仍然太粗糙:同一份 token 表示既要用于判断两个位置是否相关,又要作为相关位置被选中后传回的内容。判断相关性和传递内容不一定依赖同一组信息。一个代词位置可能需要根据语法或语义线索去匹配前面的名词,而被匹配到的名词传给后续层的内容,可以包含更完整的实体语义和上下文信息。QKV 拆分的作用,是让模型分别学习“怎么判断该读谁”和“读到以后传回什么”

这种拆分通过三组可训练线性投影完成:

Q=XWQ,K=XWK,V=XWVQ = XW_Q,\qquad K = XW_K,\qquad V = XW_V

其中 WQW_QWKW_KWVW_V 都是模型参数。可以粗略理解为:QQ(Query)表示当前位置想找什么信息,KK(Key)表示每个位置适合被什么 query 匹配到,VV(Value)表示这个位置被读取时贡献什么内容。Query 和 Key 负责判断该读谁、读多少,Value 负责定义读回来什么

如下图所示,QQKKVV 都来自同一个输入 XX,只是经过了三组不同的线性投影:

Q K V projections from one input sequence One input matrix X goes through three learned linear projections to produce Q, K, and V representations for the same tokens. X N x D W_Q D x d_k W_K D x d_k W_V D x d_v Q = XW_Q what this token looks for K = XW_K how it can be matched V = XW_V content to contribute
Q、K、V 来自同一个输入序列 X,但通过不同可训练矩阵投影到查询、匹配和内容空间。

对于单条序列,常见形状约定是:

XRN×D,Q,KRN×dk,VRN×dvX \in \mathbb{R}^{N \times D},\quad Q,K \in \mathbb{R}^{N \times d_k},\quad V \in \mathbb{R}^{N \times d_v}

这里 NN 是 token 数,DD 是输入 hidden dimension,dkd_k 是打分空间维度,dvd_v 是 value 空间维度。线性投影不会改变 token 个数,只会改变每个 token 的表示维度和用途。Q、K、V 不是三份新 token,而是同一批 token 分别用于查询、匹配和内容传递的三种表示

3. Scaled Dot-Product Attention 公式逐项推导

Scaled Dot-Product Attention 的标准公式是:

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

这行公式可以拆成四步:先用 QKTQK^T 得到所有 token 对的分数,再除以 dk\sqrt{d_k} 控制分数尺度,对每一行做 Softmax 得到读取权重,最后把权重乘到 VV 上做加权聚合。Scaled Dot-Product Attention 的本质是“打分、缩放、归一化、聚合”四步矩阵计算。如下图所示。

Scaled dot-product attention flow Q and K feed the first MatMul, then Scale, optional Mask, SoftMax, and a second MatMul with V. Scaled Dot-Product Attention Q K V MatMul Scale Mask (opt.) SoftMax MatMul
Scaled Dot-Product Attention 的四步数据流。

第一步是构造关系矩阵:

S=QKT,Sij=qikjS = QK^T,\qquad S_{ij}=q_i \cdot k_j

ii 行表示第 ii 个 token 作为 query 时,对所有 key token 的打分;第 jj 列表示第 jj 个 token 作为候选信息来源时,被其他 query 匹配的分数。SijS_{ij} 越大,表示第 ii 个位置越倾向于从第 jj 个位置读取信息。QKTQK^T 把序列内部所有位置对的匹配关系写成了一个 N×NN \times N 矩阵

第二步是缩放:

S~=Sdk\tilde{S} = \frac{S}{\sqrt{d_k}}

上一篇已经推导过,若 qqkk 的各维在简化假设下独立、均值为 0、方差为 1,则 qkq \cdot k 的方差会随 dkd_k 增大。不缩放时,维度越高,内积分数差距越容易变大,Softmax 越容易接近 one-hot 并进入梯度较弱的饱和区。除以 dk\sqrt{d_k} 的作用是控制打分尺度,而不是改变 Attention 的匹配语义

第三步是对每一行做 Softmax:

A=softmax(QKTdk)A = softmax\left(\frac{QK^T}{\sqrt{d_k}}\right)

矩阵 AA 仍然是 N×NN \times N。每一行都是一个权重分布,AijA_{ij} 表示第 ii 个位置从第 jj 个位置读取信息的比例。注意这里的权重只是当前层、当前 head 中的聚合比例,不等价于人类意义上的因果解释。Softmax 把任意实数分数转换成非负且总和为 1 的权重,因此这些权重可以用于后续的加权求和

第四步是乘以 VV

O=AVO = AV

展开到第 ii 个 token,就是:

oi=jAijvjo_i = \sum_j A_{ij}v_j

也就是说,第 ii 个输出向量不是只来自自己,而是所有 value 向量的加权和。哪些位置贡献多,由 AijA_{ij} 决定;贡献的具体内容,由 vjv_j 决定。这里的输出不是新的 token id,而是第 ii 个输入位置更新后的 hidden state。Attention 的输出是每个位置的上下文化表示:位置不变,但表示已经按权重吸收了整段序列的信息

把四步合在一起,公式里每一项的作用就清楚了:QQ 发起查询,KK 提供匹配索引,QKTQK^T 得到关系矩阵,dk\sqrt{d_k} 控制分数尺度,Softmax 得到读取权重,VV 提供被聚合内容。这些分工共同定义了一次可微的信息读取

4. Mask 约束

如果只看 softmax(QKT/dk)Vsoftmax(QK^T / \sqrt{d_k})V,每个 token 都可以读取所有位置。但真实任务经常有边界:padding 不是有效文本,原论文 decoder 的 masked self-attention 不能读取未来 target token。Mask 的作用是约束每个位置能读取哪些 token,让 Attention 符合任务规则

第一类是 padding mask。训练或批量推理时,一个 batch 里可能有不同长度的序列,短序列会被 padding 到共同长度 NN。这些 padding 位置只是为了凑齐张量形状,不应该被真实 token 读取。常见做法是在 Softmax 前把 padding 对应的 score 加上 -\infty 或一个极小值。Padding mask 解决的是“张量形状里有位置,但语义上这些位置不存在”的问题

第二类是 causal mask,也就是原论文 decoder self-attention 里的 look-ahead mask。机器翻译 decoder 生成目标序列时,第 ii 个位置只能看到自己和之前的目标 token,不能看到未来目标 token;decoder-only 语言模型的 next-token prediction 也使用同样的可见性约束。因此对第 ii 行,只允许 jij \le i 的列参与 Softmax:

允许读取的位置:

1 0 0 0
1 1 0 0
1 1 1 0
1 1 1 1

这就是下三角 mask。Causal mask 把“预测下一个 token”这个任务约束写进了 Attention 的可见范围

Mask 要放在 Softmax 前,而不是 Softmax 后简单置零。原因是 Softmax 前屏蔽非法位置后,剩余合法位置会自动重新归一化;如果先 Softmax 再置零,权重和就不再是 1,除非额外再做一次归一化。Mask 要在 Softmax 前生效,让权重只在合法位置之间重新分配

训练里还会有 loss mask,但它和 Attention mask 不是一回事。Attention mask 控制当前位置能读取哪些 token;loss mask 控制哪些位置要参与 loss 计算。比如 SFT 里,用户 prompt 通常作为上下文输入,assistant 回复部分才作为训练目标:prompt 可以被 assistant 回复读取,但 prompt 位置本身不一定计入 loss。Attention mask 管信息流,loss mask 管训练目标

5. Multi-Head Attention

单头 Attention 只有一个打分空间和一个 value 空间。语言里的关系却很多:相邻搭配、语法依赖、指代关系、长程语义呼应、格式边界都可能同时存在。如果所有关系都挤在一个 Attention 空间里,模型表达会受限。Multi-Head Attention 的动机,是让模型在多个可学习子空间里并行读取不同类型的信息

Multi-Head Attention 会把输入先投影到每个 head 的 Q,K,VQ,K,V 子空间,再把所有 head 的结果拼接起来:

MultiHead(Q,K,V)=Concat(head1,,headh)WOheadi=Attention(QWiQ,KWiK,VWiV)\begin{aligned} MultiHead(Q,K,V) &= Concat(head_1,\ldots,head_h)W^O \\ head_i &= Attention(QW_i^Q,KW_i^K,VW_i^V) \end{aligned}

每个 head 有自己的 WiQW_i^QWiKW_i^KWiVW_i^V,因此每个 head 可以学习不同的匹配方式和内容抽取方式

Multi-head attention structure Q, K, and V pass through Linear projections, several scaled dot-product attention layers run in parallel, then outputs are concatenated and projected by a final Linear layer. V K Q Linear Linear Linear Scaled Dot-Product Attention h Concat Linear
Multi-Head Attention:多个 Scaled Dot-Product Attention 层并行运行,输出经过 Concat 和 Linear 投影。

原论文的 base Transformer 使用 dmodel=512d_{\text{model}}=512h=8h=8,因此每个 head 的 dk=dv=64d_k=d_v=64。这样做不是把总维度放大 8 倍,而是把 512 维表示切成 8 个并行子空间,再用输出投影混合回 dmodeld_{\text{model}}论文里的 Multi-Head Attention 既增加了表示子空间数量,又把每个 head 的计算维度控制在较小范围内

形状上,常见设置是 D=hdheadD = h \cdot d_{\text{head}}。实现时通常把 B×N×DB \times N \times D reshape 成包含 head 维度的张量,让每个 head 在 dheadd_{\text{head}} 维子空间里做 Attention,再把结果拼回 DD 维。图中的 Linear、并行 heads、Concat 和输出投影,对应的就是这个张量组织过程。Multi-Head Attention 首先是表达能力设计,同时也天然适合 GPU 批量矩阵计算

不同 head 学到了什么,不应该被过度解释成确定的人类规则。某些 head 可能更关注局部位置,某些 head 可能更关注分隔符或长程依赖,但这不是架构显式规定的语义分工,而是训练目标和数据共同塑造的结果。Attention head 可以形成有用的读取模式,但单个 head 的权重不等于稳定的自然语言解释

6. 位置编码与 RoPE

前文主要以 self-attention 为例。Self-attention 指 Q,K,VQ,K,V 都由同一段 token 序列投影得到:Q=XWQQ=XW^QK=XWKK=XW^KV=XWVV=XW^V;encoder self-attention、decoder masked self-attention 都属于这个范畴,而 cross-attention 则通常是 QQ 来自 decoder、K,VK,V 来自 encoder。Self-attention 是 Attention 算子的一个使用场景,不是另一套不同的公式

这里的 Attention 算子指 softmax(QKT/dk)Vsoftmax(QK^T/\sqrt{d_k})V 这套打分、缩放、归一化和加权求和计算。它接收的是 Q,K,VQ,K,V 的数值矩阵;矩阵第几行在程序里当然对应第几个 token,但这个行号没有作为数值特征进入 QKTQK^T行顺序会被计算保留下来,但“第几个位置”和“相隔多远”不会自动变成模型可使用的信息

如果用置换矩阵 PP 同时重排输入行,纯 self-attention 的输出也会按同样方式重排:SelfAttention(PX)=PSelfAttention(X)SelfAttention(PX)=P\,SelfAttention(X)。这说明它对行排列是置换等变的:程序知道行的顺序,算子本身只根据行向量之间的匹配计算权重。位置机制的作用,是把外部行号或相对距离转成数值信号,注入 token 表示或 Attention 打分过程

原始 Attention Is All You Need 使用正弦位置编码:为第 tt 个位置生成位置向量 ptp_t,再与 token embedding ete_t 相加:

xt=et+ptx_t = e_t + p_t

这样,同一个 token 出现在不同位置时,进入后续层的表示也不同。绝对位置编码的核心做法,是把“这个 token 在哪里”直接写进输入表示里

ptp_t 是一个 dmodeld_{\text{model}} 维向量;下面的 PE(pos,2j)PE(pos,2j)PE(pos,2j+1)PE(pos,2j+1) 只是这个向量在第 2j2j 和第 2j+12j+1 个维度上的两个标量分量:

PE(pos,2j)=sin(pos100002j/dmodel)PE(pos,2j+1)=cos(pos100002j/dmodel)\begin{aligned} PE(pos,2j) &= \sin\left(\frac{pos}{10000^{2j/d_{\text{model}}}}\right) \\ PE(pos,2j+1) &= \cos\left(\frac{pos}{10000^{2j/d_{\text{model}}}}\right) \end{aligned}

把所有维度的标量分量排在一起,才得到完整的位置向量 pposp_{pos}。这里 pospos 是位置,jj 是维度对索引;偶数维用正弦,奇数维用余弦,不同维度对应不同波长。因为这个函数可以在任意 pospos 上计算,模型不需要为每个最大位置单独学习参数;论文也希望这种形式更容易外推到训练长度之外。正弦位置编码的关键信息,是用固定函数把位置变成可加到 token 表示上的向量

现代 decoder-only LLM 中经常见到 RoPE(Rotary Position Embedding)。它不是原始 Transformer 论文里的位置编码,但很适合和 Attention 公式放在一起理解,因为它作用在 QQKK 上,而不是简单加到输入 embedding 上。二维旋转可以写成:

Rθ=[cosθsinθsinθcosθ]R_\theta = \begin{bmatrix} \cos\theta & -\sin\theta \\ \sin\theta & \cos\theta \end{bmatrix}

RoPE 会分别在每个 query 向量和 key 向量内部,把相邻两个维度分成一组二维平面:(0,1)(0,1)(2,3)(2,3)、……。第 mm 个位置的 query 按角度 mθm\theta 旋转这些二维分量,第 nn 个位置的 key 按角度 nθn\theta 旋转这些二维分量。两者做内积时,共同的绝对位置部分会抵消,真正影响位置关系的是角度差 (nm)θ(n-m)\thetaRoPE 的关键直觉是:把位置写成旋转角度后,Attention 打分会自然包含 query 和 key 的相对距离

长上下文会直接触及这个机制:当序列长度超过训练时常见范围,query 和 key 的相对距离变大,对应的旋转角度也进入模型较少见过的区间。Attention 的打分来自 QKTQK^T,RoPE 又直接改变 QQKK 的几何关系,所以位置外推、上下文扩展和 KV Cache 讨论经常会回到 RoPE。

7. Transformer Block 中的 Attention

Transformer encoder-decoder architecture Encoder and decoder Transformer architecture with embeddings, positional encoding, attention, feed-forward, add and norm layers, linear, softmax, and output probabilities. Add & Norm Feed Forward Add & Norm Multi-Head Attention Input Embedding Add & Norm Feed Forward Add & Norm Multi-Head Attention Add & Norm Masked Multi-Head Attention Output Embedding Linear Softmax N x N x Positional Encoding Positional Encoding Inputs Outputs (shifted right) Output Probabilities
原论文 Figure 1 的 encoder-decoder Transformer 结构:encoder 处理 source tokens,decoder 根据右移一位后的目标序列生成输出。

原论文的 Transformer 是 sequence-to-sequence 结构。Encoder 读取 source tokens,输出一组上下文表示;decoder 读取右移一位后的目标序列,并通过 causal mask 保证每个位置只能看到当前位置及之前的 token,再通过 encoder-decoder attention 读取 encoder 输出。Encoder 负责把输入序列编码成可被读取的表示,decoder 负责在已生成前缀和输入表示的条件下生成输出序列

图里的 Add & Norm 对应残差连接加层归一化。对一个子层 SublayerSublayer,原论文使用 post-LN 写法:Add & Norm 先把子层输出加回原表示,再对加和结果做归一化

LayerNorm(x+Sublayer(x))LayerNorm(x + Sublayer(x))

这里的 Add 是把子层输出当作修正量加回原表示,Norm 再把加和后的 hidden vector 归一化。Add & Norm 的作用,是让每个子层在保留原表示的基础上做更新,并让多层堆叠时的数值尺度更稳定

Position-wise feed-forward network 对每个位置独立应用同一个两层 MLP:它是逐 token 的非线性变换,不负责不同位置之间的信息交换

FFN(x)=max(0,xW1+b1)W2+b2FFN(x)=\max(0,xW_1+b_1)W_2+b_2

在 base Transformer 中,dmodel=512d_{\text{model}}=512,中间层维度 dff=2048d_{ff}=2048。它不直接混合不同 token,而是在每个 token 自己的 hidden vector 上增加非线性变换能力。Attention 负责跨位置通信,position-wise FFN 负责逐位置处理已经聚合到的表示

图顶端的 Linear + Softmax 是 decoder 的输出层。Decoder 最后一层的 hidden vector 先经过线性投影得到词表上每个 token 的 logit,再经 Softmax 变成下一个 token 的概率分布。Embedding 和位置编码把离散 token 变成输入表示,Linear + Softmax 把输出表示变回词表分布

原论文选择 self-attention 放进 encoder 和 decoder 的一个核心理由是计算路径更短、并行度更高。用 nn 表示序列长度、dd 表示表示维度时,self-attention 每层复杂度是 O(n2d)O(n^2d),但顺序操作数是 O(1)O(1),任意两个位置之间的最大路径长度也是 O(1)O(1);RNN 的顺序操作数和最大路径长度都是 O(n)O(n)Attention 的优势不只是“能看全局”,还在于它把序列依赖从时间递推改成了可并行的矩阵计算

Layer typeComplexity per layerSequential operationsMaximum path length
Self-attentionO(n2d)O(n^2d)O(1)O(1)O(1)O(1)
RecurrentO(nd2)O(nd^2)O(n)O(n)O(n)O(n)
ConvolutionalO(knd2)O(knd^2)O(1)O(1)O(logkn)O(\log_k n)

8. 自回归推理与 KV Cache

KV Cache 不是原论文提出的模型结构,而是同一 Attention 算子在现代自回归推理中的实现技术。训练时可以并行处理整段序列,但推理时语言模型要自回归生成:先根据已有上下文预测下一个 token,再把新 token 接回上下文继续预测。推理因此通常分成 prefill 和 decode 两段。KV Cache 的价值,只在理解 prefill/decode 差异之后才真正清楚

Prefill 阶段处理用户已经输入的 prompt。模型把整段 prompt 一次送入各层 Attention;在 causal mask 下,每个位置只能读取自己和之前的位置,同时每一层都会为 prompt 的每个 token 计算并保存 K,VK,V。预测第一个新 token 时,最后一个 prompt 位置的输出给出 logits;这个输出已经由最后一个 prompt 位置的 QQ 对 prompt 内可见的 K,VK,V 做过 Attention。Prefill 的作用是先处理已有上下文,并建立后续 decode 要反复读取的历史 K,VK,V 缓存

Decode 阶段从已经生成出第一个新 token 之后开始。每一步只把上一步生成的 token 作为当前输入,计算它在各层里的 Q,K,VQ,K,V;新的 K,VK,V 追加进同一个请求的 KV Cache,新的 QQ 和从 prompt 到当前 token 的所有 KK 做匹配,再按权重聚合对应的 VVDecode 不重新计算历史 token,只让当前 token 的 Query 读取同一个请求已经缓存下来的上下文

下图只画 decode 的单步读写关系。当前 token 计算自己的 Q,K,VQ,K,V,其中 K,VK,V 写入缓存,QQ 用来对缓存中的 K,VK,V 做 Attention

KV Cache during decode A new token computes Q K V, appends K and V into the cache, and uses Q to attend over cached K and V. new token current step Q_new K_new V_new KV Cache history K / V for this request K_1 ... K_t, K_new V_1 ... V_t, V_new Q_new reads cached K/V append K append V query path
Decode 阶段的 KV Cache:当前 token 产生 Q/K/V,新的 K/V 追加到缓存,当前 Q 对同一请求的历史 K/V 做 Attention。

如果没有 KV Cache,第 tt 步生成时就要重新计算前 t1t-1 个历史 token 的 key 和 value。可是模型参数固定、历史 token 表示在给定上下文下已经算过,它们的 KKVV 不需要每步从头再来。KV Cache 避免的是自回归 decode 中对历史 Key/Value 的重复计算

每进入一个 decode 步,KV Cache 的机制可以写成三步:

  1. 对当前输入 token 计算 Q,K,VQ,K,V
  2. 把当前 token 的 K,VK,V 追加到这个请求自己的 KV Cache。
  3. 用当前 token 的 QQ 和缓存里的所有 KK 计算 Attention 权重,再用这些权重从对应的 VV 中汇总信息。

这是一种空间换时间:用 HBM 保存历史 key/value,换取后续步骤少做重复矩阵投影和历史计算。KV Cache 让 decode 更快,但代价是显存占用随 batch、层数和上下文长度增长

粗略估算 KV Cache 大小,可以看这些因子:

KV Cache2×B×L×N×hkv×dhead×bytesKV\ Cache \approx 2 \times B \times L \times N \times h_{kv} \times d_{\text{head}} \times bytes

其中 2 表示 K 和 V,BB 是 batch size,LL 是层数,NN 是缓存 token 数,hkvh_{kv} 是 KV head 数,dheadd_{\text{head}} 是每个 head 的维度,bytesbytes 是数据类型字节数。KV Cache 的成本不是一个常数,而是随请求并发、上下文长度和模型结构一起放大

这也解释了为什么推理服务经常受显存容量和带宽约束。长上下文、多并发、层数多、head 数多都会放大 KV Cache;decode 每步虽然只新增一个 token,却要不断读取历史 K/V。自回归推理的瓶颈常常不是“当前 token 算不算得动”,而是“历史缓存放不放得下、读不读得快”

9. KV Cache 压缩与内存优化

这一节讨论的 MQA、GQA、FlashAttention 和 PagedAttention 都不是 2017 年原论文里的内容,而是 Attention 在现代长上下文和在线推理中遇到内存瓶颈后的工程演进。标准 Multi-Head Attention 中,每个 query head 都有自己的 K/V head;如果 head 很多,KV Cache 也会很大。MQA(Multi-Query Attention)让多个 query head 共享一套 K/V,GQA(Grouped-Query Attention)则把 query heads 分组,每组共享一套 K/V。MQA/GQA 的核心目标,是在模型质量、KV Cache 大小和推理吞吐之间做折中

hh 表示 query head 数,hkvh_{kv} 表示 KV head 数。标准 MHA 通常有 hkv=hh_{kv}=h;MQA 接近 hkv=1h_{kv}=1;GQA 介于两者之间。hkvh_{kv} 越小,KV Cache 公式中的对应因子越小,显存和带宽压力也越低,但共享 K/V 可能带来表达能力损失。GQA 之所以常见,是因为它比 MQA 保留更多表示能力,又比标准 MHA 更省 KV Cache

FlashAttention 解决的是另一类问题。标准 Attention 的朴素实现会产生 N×NN \times N 的 score 矩阵和 probability 矩阵,长序列下这些中间结果会带来大量 HBM 读写。FlashAttention 不改变 Attention 的数学结果,而是把 QQKKVV 分块搬到更快的片上存储附近,使用 tiling 和 online softmax 避免把完整 attention 矩阵写回 HBM。FlashAttention 优化的是数据搬运路径,不是把 Attention 近似成另一个算法

Online softmax 的直觉是:Softmax 需要一行的最大值和归一化分母,但分块计算时不能一次看到整行。FlashAttention 维护每一行当前看到的最大值和归一化统计量,在新 block 到来时更新这些量,并把之前的局部结果按比例修正。它让分块计算仍然得到与完整 Softmax 一致的结果,同时减少 HBM 往返

PagedAttention 关注的是推理服务里的 KV Cache 管理。在线服务中,不同请求长度不同、到达时间不同、结束时间不同;如果给每个请求预留连续大块显存,容易出现碎片和浪费。PagedAttention 借鉴虚拟内存分页思想,把 KV Cache 切成固定大小 block,用 block table 维护逻辑位置到物理显存块的映射。FlashAttention 优化单次 Attention 计算的 I/O,PagedAttention 优化高并发推理中的 KV Cache 内存管理

这些技术看起来分散,但都围绕同一件事:让有限的 GPU 存储层级服务更多有效 token。MHA/GQA 改变 KV head 数,FlashAttention 改变中间矩阵的读写方式,PagedAttention 改变 KV Cache 的分配和复用方式。现代 Attention 优化的主线,是在不改变核心语义的前提下减少重复计算、显存占用和 HBM 往返

结语

现在可以把 Attention 分成三层理解。数学层里,QKTQK^T 负责打分,dk\sqrt{d_k} 控制尺度,Softmax 归一化成读取权重,VV 提供被聚合内容。Scaled Dot-Product Attention 是一次可微的信息读取,而不是一串孤立的矩阵符号

架构层里,mask 约束可见范围,Multi-Head Attention 提供多个匹配子空间,位置机制注入顺序信息,Transformer block 用 Residual Connection、LayerNorm 和 MLP 把 Attention 稳定堆叠成深层网络。Attention 是核心算子,但 Transformer 的能力来自 Attention 与其他子层的组合

系统层里,KV Cache 支撑自回归推理,MQA/GQA 降低 KV 成本,FlashAttention 减少 HBM 往返,PagedAttention 改善高并发场景下的显存管理。大模型里的 Attention 既要在数学上表达依赖关系,也要在工程上适应 GPU 存储层级和在线推理负载