AI 工程基础体系 · 第 12/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

Transformer 完整原理:Attention、残差、归一化、解码和 KV Cache

Transformer 是一种以注意力机制为核心的序列建模架构。它不依赖循环神经网络逐步传递隐藏状态,而是通过矩阵运算让一个位置直接读取其他位置的信息。今天的大语言模型通常采用 Transformer 的 decoder-only 变体,但“Transformer”这个名称最初指的是论文《Attention Is All You Need》中提出的 encoder-decoder 架构。

要真正理解大语言模型的训练和推理,不能只记住“Transformer 使用 Attention”。至少需要把以下链路连起来:

文本
  ↓
Tokenization
  ↓
Token ID
  ↓
Embedding + 位置表示
  ↓
Transformer Block
  ├─ Causal Self-Attention
  ├─ Residual Connection
  ├─ Normalization
  └─ Feed-Forward Network
  ↓
隐藏状态
  ↓
LM Head
  ↓
Logits
  ↓
Softmax / 采样 / 贪心解码
  ↓
下一个 Token

其中,训练通常一次处理完整序列;生成式推理则需要反复预测下一个 Token。KV Cache 正是为了避免生成过程中反复计算历史 Token 的 Key 和 Value。


一、从 Token 到隐藏状态

1. Tokenization:模型处理的是 Token,不是字符或单词

Tokenization 是把文本转换为离散 Token 序列的过程。例如:

输入:
Transformer 很强大

可能得到:
["Transformer", " 很", "强", "大"]

实际切分结果取决于词表、训练语料和 Tokenizer 算法。常见方法包括:

  • BPE:Byte Pair Encoding;
  • WordPiece;
  • Unigram;
  • byte-level BPE;
  • 基于字节或字符的变体。

模型真正接收的是整数 ID:

[18432, 912, 3071, 58]

Token ID 本身没有大小关系。ID 为 1000 并不比 ID 为 500 更“接近”某个词。它只是词表中的索引。

设词表大小为 VV,输入序列长度为 nn,则 Token ID 序列为:

x1,x2,,xn,xi{0,1,,V1}x_1, x_2, \dots, x_n,\quad x_i \in \{0,1,\dots,V-1\}

1.1 Embedding:把离散索引映射到向量

Embedding 矩阵记为:

ERV×dmodelE \in \mathbb{R}^{V \times d_{\text{model}}}

其中:

  • VV:词表大小;
  • dmodeld_{\text{model}}:模型隐藏维度;
  • EiE_i:第 ii 个 Token 的向量。

对 Token ID xix_i,Embedding 操作就是取出对应行:

ei=E[xi]e_i = E[x_i]

因此,输入序列经过 Embedding 后得到:

XRn×dmodelX \in \mathbb{R}^{n \times d_{\text{model}}}

Embedding 的训练目标不是显式地让“同义词距离接近”,而是在语言建模损失驱动下学习有用的表示。相近语义、相似句法或相似上下文中的 Token,可能逐渐形成有结构的向量空间,但这不是由 Embedding 层单独保证的。

1.2 位置表示:Attention 本身不认识顺序

如果只对输入 Token 做自注意力,输入:

猫追狗

和输入:

狗追猫

只要 Token 集合相同,注意力机制本身无法区分它们的顺序。因为标准 Self-Attention 对输入行进行置换时,也会同步置换输出,缺少额外的位置来源。

因此需要位置表示。常见方式有:

  1. 绝对位置 Embedding
    为位置 0,1,2,0,1,2,\dots 学习一个位置向量:

    hi(0)=ei+pih_i^{(0)} = e_i + p_i

  2. 正弦位置编码
    原始 Transformer 使用固定函数:

    PE(pos,2k)=sin(pos100002k/dmodel)PE(pos,2k)=\sin\left(\frac{pos}{10000^{2k/d_{\text{model}}}}\right)

    PE(pos,2k+1)=cos(pos100002k/dmodel)PE(pos,2k+1)=\cos\left(\frac{pos}{10000^{2k/d_{\text{model}}}}\right)

  3. 相对位置表示
    让注意力分数依赖两个位置之间的相对距离。

  4. RoPE,即 Rotary Position Embedding
    将 Query 和 Key 的二维子空间进行与位置相关的旋转,使点积携带相对位置信息。许多现代 decoder-only 模型使用这一类方法,但具体实现、缩放方式和长上下文扩展策略取决于模型。

位置编码不是“给每个词附加一个坐标”这么简单。它必须让后续的点积、卷积式局部模式或其他变换能够识别顺序和距离。


二、Self-Attention 的形式化定义

2.1 Query、Key、Value 分别是什么

给定输入矩阵:

XRn×dmodelX \in \mathbb{R}^{n \times d_{\text{model}}}

通过三个不同的线性变换得到:

Q=XWQQ = XW_Q

K=XWKK = XW_K

V=XWVV = XW_V

其中:

  • QQ:Query,当前位置想查询什么;
  • KK:Key,每个位置可以被如何匹配;
  • VV:Value,被读取的实际内容;
  • WQ,WK,WVW_Q,W_K,W_V:可训练矩阵。

假设单头注意力维度为 dkd_k,则:

WQ,WKRdmodel×dkW_Q,W_K \in \mathbb{R}^{d_{\text{model}}\times d_k}

通常:

Q,K,VRn×dkQ,K,V \in \mathbb{R}^{n\times d_k}

对位置 ii 和位置 jj,匹配分数是:

sij=qikjdks_{ij} = \frac{q_i k_j^\top}{\sqrt{d_k}}

分母 dk\sqrt{d_k} 是缩放因子。若 Query 和 Key 的每个维度近似独立、方差为常数,则点积的方差会随 dkd_k 增长。维度较大时,不缩放的点积会让 Softmax 进入极端饱和区,梯度变得很小。

对固定的查询位置 ii,经过 Softmax 后:

aij=exp(sij)t=1nexp(sit)a_{ij} = \frac{\exp(s_{ij})} {\sum_{t=1}^{n}\exp(s_{it})}

注意力输出为:

oi=j=1naijvjo_i = \sum_{j=1}^{n}a_{ij}v_j

矩阵形式为:

Attention(Q,K,V)=softmax(QKdk)V\operatorname{Attention}(Q,K,V) = \operatorname{softmax} \left( \frac{QK^\top}{\sqrt{d_k}} \right)V

这里的 Softmax 是按每一行计算,因此每个查询位置对所有 Key 的权重和为 1。输出是 Value 的加权和,而不是 Query 和 Key 的简单平均。

2.2 完整数值算例

设有 2 个位置,且:

Q=[1001],K=[1011],V=[100020]Q= \begin{bmatrix} 1&0\\ 0&1 \end{bmatrix}, \quad K= \begin{bmatrix} 1&0\\ 1&1 \end{bmatrix}, \quad V= \begin{bmatrix} 10&0\\ 0&20 \end{bmatrix}

这里 dk=2d_k=2,所以缩放因子为 2\sqrt{2}

先计算 QKQK^\top

QK=[1101]QK^\top = \begin{bmatrix} 1&1\\ 0&1 \end{bmatrix}

缩放后:

S=[1/21/201/2][0.7070.70700.707]S= \begin{bmatrix} 1/\sqrt{2}&1/\sqrt{2}\\ 0&1/\sqrt{2} \end{bmatrix} \approx \begin{bmatrix} 0.707&0.707\\ 0&0.707 \end{bmatrix}

对第一行做 Softmax,因为两个值相同:

softmax(0.707,0.707)=(0.5,0.5)\operatorname{softmax}(0.707,0.707) = (0.5,0.5)

对第二行:

softmax(0,0.707)(0.330,0.670)\operatorname{softmax}(0,0.707) \approx (0.330,0.670)

因此注意力矩阵近似为:

A[0.50.50.3300.670]A\approx \begin{bmatrix} 0.5&0.5\\ 0.330&0.670 \end{bmatrix}

最后:

AV=[0.50.50.3300.670][100020]AV = \begin{bmatrix} 0.5&0.5\\ 0.330&0.670 \end{bmatrix} \begin{bmatrix} 10&0\\ 0&20 \end{bmatrix}

得到:

O[5103.313.4]O\approx \begin{bmatrix} 5&10\\ 3.3&13.4 \end{bmatrix}

第一个位置对两个 Value 各读取一半;第二个位置更偏向第二个 Value。

2.3 为什么需要三个投影,而不是直接对输入做相似度

如果令 Q=K=V=XQ=K=V=X,就只能使用输入向量自身的相似度。实际模型需要把同一个隐藏状态投影到不同的功能空间:

  • WQW_Q 学习“我当前需要什么信息”;
  • WKW_K 学习“我能以什么方式被匹配”;
  • WVW_V 学习“被匹配后应该提供什么内容”。

例如在代词消解中,当前位置的 Query 可能寻找“与当前主语一致的实体”,而历史实体的 Key 暴露其句法或语义特征,Value 则包含可供后续层使用的表示。


三、Causal Mask:生成模型为什么不能偷看未来

训练自回归语言模型时,目标是预测:

P(xtx1,,xt1)P(x_t\mid x_1,\dots,x_{t-1})

因此位置 tt 不能读取 xt+1x_{t+1} 及之后的 Token。实现上,在 Softmax 前对未来位置加上 -\infty

Mij={0,ji,j>iM_{ij}= \begin{cases} 0,&j\le i\\ -\infty,&j>i \end{cases}

注意力变为:

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

例如,长度为 3 的因果掩码为:

M=[000000]M= \begin{bmatrix} 0&-\infty&-\infty\\ 0&0&-\infty\\ 0&0&0 \end{bmatrix}

加掩码后,第一行 Softmax 只会在第一个位置分配概率,第二行只会在前两个位置分配概率。

这也解释了为什么训练可以并行:虽然每个位置只能看过去,但所有位置的 Query、Key、Value 可以一次性算出,矩阵乘法和带掩码的 Softmax 同时处理整段序列。

3.1 反例:没有 Causal Mask 会发生什么

假设训练目标是:

输入:我 喜欢
标签:喜欢 苹果

如果第二个位置在计算“喜欢”对应的预测时可以读取未来的“苹果”,模型就可能直接复制未来 Token,而不是学习语言规律。训练损失会异常地低,但部署时生成阶段没有未来 Token,性能会突然崩溃。这种情况称为信息泄漏。

Padding Mask 是另一种掩码。它用于阻止模型读取批次中补齐的 <PAD> 位置。Causal Mask 解决时间方向上的未来泄漏,Padding Mask 解决批处理中的无效位置,两者用途不同,实际实现可能需要组合。


四、多头注意力:在多个子空间中同时匹配

单个注意力头只有一组 WQ,WK,WVW_Q,W_K,W_V。多头注意力将隐藏维度划分为 hh 个头:

headr=Attention(QWQ(r),KWK(r),VWV(r))\operatorname{head}_r = \operatorname{Attention} (QW_Q^{(r)},KW_K^{(r)},VW_V^{(r)})

然后拼接所有头:

MHA(X)=Concat(head1,,headh)WO\operatorname{MHA}(X) = \operatorname{Concat}(\operatorname{head}_1,\dots,\operatorname{head}_h)W_O

如果:

dmodel=4096,h=32d_{\text{model}}=4096,\quad h=32

常见情况下每个头的维度为:

dk=dv=4096/32=128d_k=d_v=4096/32=128

不同头可以学习不同关系,例如局部邻近、长距离指代、句法依赖或格式结构。但“某个头一定对应某种人类可解释功能”并不是规范保证,头的功能可能重叠,也可能随层和模型训练而变化。

4.1 MHA、MQA 和 GQA

标准 Multi-Head Attention,简称 MHA,通常每个头都有独立的 Query、Key、Value:

Query heads: h
Key heads:   h
Value heads: h

Multi-Query Attention,简称 MQA,让多个 Query 头共享一组 Key 和 Value:

Query heads: h
Key heads:   1
Value heads: 1

Grouped-Query Attention,简称 GQA,则让多个 Query 头分组共享 Key 和 Value:

Query heads: h
Key heads:   g
Value heads:  g

其中 1<g<h1<g<h

这三者的主要区别在推理时的 KV Cache 大小。Query 通常只需要当前 Token 的计算,而历史 Key 和 Value 需要保存。因此减少 KV 头数可以显著降低缓存占用和内存带宽压力,但可能改变模型容量和效果,不能无损地假设为“只是优化”。


五、Transformer Block:Attention 不是完整的 Transformer

一个 Transformer Block 通常由以下部分组成:

  1. 归一化;
  2. 自注意力;
  3. 残差连接;
  4. 归一化;
  5. 前馈网络;
  6. 残差连接。

原始论文使用 Post-LN 结构,现代大语言模型常使用 Pre-LN 或 RMSNorm 变体。

5.1 残差连接:保留旧表示并叠加新变换

残差连接的基本形式是:

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

其中 FF 可以是注意力子层或前馈子层。

在一个 Block 中,典型 Pre-LN 形式是:

u=x+Attention(Norm(x))u=x+\operatorname{Attention}(\operatorname{Norm}(x))

y=u+FFN(Norm(u))y=u+\operatorname{FFN}(\operatorname{Norm}(u))

残差连接有两个重要作用。

第一,它允许每个子层学习“对已有表示进行增量修改”,而不是每层都必须重新构造完整表示。若某个子层接近零,信息仍可沿恒等路径传递。

第二,它改善深层网络的梯度传播。设:

yl=xl+Fl(xl)y_l=x_l+F_l(x_l)

则:

ylxl=I+Flxl\frac{\partial y_l}{\partial x_l} = I+\frac{\partial F_l}{\partial x_l}

梯度中始终存在恒等矩阵 II 这一条路径。它不能保证训练必然稳定,但比单纯连续复合变换更容易优化。

5.2 残差不是“复制一份输入”

残差连接的输出维度必须与输入一致:

x,F(x)Rdmodelx,F(x)\in\mathbb{R}^{d_{\text{model}}}

如果子层内部使用不同维度,例如 FFN 中间维度为 dffd_{\text{ff}},最终必须投影回 dmodeld_{\text{model}},才能与输入相加。

残差也不等于把所有层的输出简单相加。每个子层有自己的归一化、投影和非线性,且通常按顺序更新状态:

x0
 ↓
x1 = x0 + Attention(Norm(x0))
 ↓
x2 = x1 + FFN(Norm(x1))

六、归一化:为什么要稳定每个 Token 的通道尺度

6.1 LayerNorm

对单个位置的隐藏向量:

x=(x1,,xd)x=(x_1,\dots,x_d)

LayerNorm 计算均值和方差:

μ=1dj=1dxj\mu=\frac{1}{d}\sum_{j=1}^{d}x_j

σ2=1dj=1d(xjμ)2\sigma^2=\frac{1}{d}\sum_{j=1}^{d}(x_j-\mu)^2

然后:

LayerNorm(x)j=γjxjμσ2+ϵ+βj\operatorname{LayerNorm}(x)_j = \gamma_j \frac{x_j-\mu}{\sqrt{\sigma^2+\epsilon}} +\beta_j

其中:

  • γ,β\gamma,\beta:可训练的逐通道缩放和偏置;
  • ϵ\epsilon:防止除零的小常数;
  • 归一化维度通常是隐藏维度 dd,不是 Token 序列维度。

LayerNorm 不会把整个 Batch 中不同样本混在一起,因此适合变长序列和小 Batch 推理。它与 BatchNorm 的统计方式不同,不能互换。

6.2 RMSNorm

RMSNorm 不减去均值,只使用均方根:

RMS(x)=1dj=1dxj2+ϵ\operatorname{RMS}(x) = \sqrt{\frac{1}{d}\sum_{j=1}^{d}x_j^2+\epsilon}

RMSNorm(x)j=gjxjRMS(x)\operatorname{RMSNorm}(x)_j = g_j\frac{x_j}{\operatorname{RMS}(x)}

它计算较简单,在许多现代语言模型中使用。RMSNorm 没有 LayerNorm 的中心化步骤,所以两者不是同一个公式,也不能在加载模型权重时随意替换。

6.3 Pre-LN 与 Post-LN

Post-LN 的原始形式近似为:

y=Norm(x+Attention(x))y=\operatorname{Norm}(x+\operatorname{Attention}(x))

z=Norm(y+FFN(y))z=\operatorname{Norm}(y+\operatorname{FFN}(y))

Pre-LN 的形式是:

y=x+Attention(Norm(x))y=x+\operatorname{Attention}(\operatorname{Norm}(x))

z=y+FFN(Norm(y))z=y+\operatorname{FFN}(\operatorname{Norm}(y))

Pre-LN 的残差路径更直接,深层训练通常更容易;但这不是说 Post-LN 一定错误。模型的层数、初始化、学习率、归一化类型和训练策略都影响稳定性。加载已有模型时必须保持原始架构的归一化位置,否则即使张量形状相同,数值行为也会改变。


七、前馈网络:Attention 后的逐位置非线性变换

Attention 负责位置之间的信息交换,Feed-Forward Network,简称 FFN,负责对每个位置独立进行通道变换。

经典形式是:

FFN(x)=W2σ(W1x+b1)+b2\operatorname{FFN}(x) = W_2\,\sigma(W_1x+b_1)+b_2

其中:

  • W1W_1 把维度从 dmodeld_{\text{model}} 扩展到 dffd_{\text{ff}}
  • σ\sigma 是激活函数,如 ReLU 或 GELU;
  • W2W_2 再投影回 dmodeld_{\text{model}}

对序列矩阵 XX,FFN 是逐行应用的:

FFN(X)i=FFN(Xi)\operatorname{FFN}(X)_i=\operatorname{FFN}(X_i)

它不直接混合不同 Token 的位置。位置之间的信息混合主要发生在 Attention 中。

现代模型常见 SwiGLU 一类门控 FFN:

FFN(x)=W2(SiLU(Wgx)Wupx)\operatorname{FFN}(x) = W_2\left(\operatorname{SiLU}(W_gx)\odot W_upx\right)

其中 \odot 是逐元素乘法。门控分支使模型可以根据输入控制哪些通道被放大。使用 SwiGLU 时,为了控制总参数量,dffd_{\text{ff}} 往往不会简单地等于经典 FFN 的同一数值。


八、从隐藏状态到 Logits

经过 LL 层 Transformer Block 后,得到:

HRn×dmodelH\in\mathbb{R}^{n\times d_{\text{model}}}

对 decoder-only 语言模型,通常取最后一个位置的隐藏状态 hth_t,经过语言模型头:

zt=htWlm+bz_t=h_tW_{\text{lm}}+b

其中:

ztRVz_t\in\mathbb{R}^{V}

ztz_t 称为 logits,是每个词表 Token 的未归一化分数。

概率为:

P(xt+1=vxt)=exp(zt,v)u=1Vexp(zt,u)P(x_{t+1}=v\mid x_{\le t}) = \frac{\exp(z_{t,v})} {\sum_{u=1}^{V}\exp(z_{t,u})}

训练时通常使用交叉熵:

L=t=1n1logP(xt+1xt)\mathcal{L} = -\sum_{t=1}^{n-1} \log P(x_{t+1}\mid x_{\le t})

实际实现常将 logits 与标签错位:

输入位置:  x0  x1  x2  x3
预测目标:  x1  x2  x3  x4

如果模型使用权重绑定,语言模型头可能复用输入 Embedding 矩阵的转置:

Wlm=EW_{\text{lm}}=E^\top

但权重绑定是常见设计,不是 Transformer 架构的强制规范。


九、解码:模型如何从概率得到文本

模型每一步只输出一个 Token 分布。解码器根据这个分布选择 Token,然后把新 Token 放回输入,继续预测。

9.1 贪心解码

选择概率最高的 Token:

xt+1=argmaxvP(vxt)x_{t+1}=\arg\max_v P(v\mid x_{\le t})

优点是简单、稳定、可复现;缺点是局部最优可能导致重复或缺乏多样性。

9.2 Temperature

温度 T>0T>0 作用于 logits:

PT(v)=softmax(z/T)vP_T(v) = \operatorname{softmax}(z/T)_v

  • T<1T<1:分布更尖锐,偏向高概率 Token;
  • T>1T>1:分布更平坦,随机性增加;
  • T0T\to 0:趋近贪心选择,但实际实现通常需要避免直接除以零。

Temperature 不会改变 Token 的排序,只改变概率差异。

9.3 Top-k 和 Top-p

Top-k 只保留概率最高的 kk 个候选,其余 logits 设为 -\infty

Top-p,也称 nucleus sampling,按概率从高到低累加,只保留累计概率达到 pp 的最小候选集合。它的候选数量会随分布变化:

  • 分布尖锐时,候选较少;
  • 分布平坦时,候选较多。

过滤后必须重新归一化概率,不能在原分布上直接把剩余概率当作最终概率。

9.4 Beam Search 与采样的边界

Beam Search 保留多个累计对数概率最高的序列,适合某些机器翻译、结构化生成任务,但不等价于“更聪明的随机采样”。对开放式对话,Beam Search 可能产生重复、模板化输出。

9.5 EOS、最大长度和停止条件

解码通常在以下条件之一满足时停止:

  • 生成了 EOS Token;
  • 达到最大新 Token 数;
  • 命中业务停止字符串;
  • 结构化输出解析成功;
  • 超时或预算耗尽。

停止字符串需要注意 Token 边界和字节边界。简单地对每个 Token 的字符串片段做匹配,可能无法正确处理跨 Token 的停止序列。


十、Decoder-only 与 Encoder-Decoder 的区别

10.1 Decoder-only

GPT 类模型使用 decoder-only 结构。所有层主要执行带 Causal Mask 的 Self-Attention:

[输入 Token] → Causal Self-Attention → FFN → ... → 下一个 Token

它适合统一的下一 Token 预测,也适合通过提示词完成问答、代码和文本生成。

10.2 原始 Transformer 的 Encoder-Decoder

原始论文中的 Transformer 包含:

  • Encoder:双向 Self-Attention,可读取输入序列全部位置;
  • Decoder:带 Causal Mask 的 Self-Attention;
  • Cross-Attention:Decoder 的 Query 读取 Encoder 输出的 Key 和 Value。

Cross-Attention 形式为:

Attention(Qdec,Kenc,Venc)\operatorname{Attention}(Q_{\text{dec}},K_{\text{enc}},V_{\text{enc}})

这里:

  • Query 来自 Decoder 当前隐藏状态;
  • Key 和 Value 来自 Encoder 输出;
  • Decoder 不能读取未来目标 Token,但可以读取完整源输入。

机器翻译就是典型例子:

源语言 → Encoder → 编码表示
                         ↑
目标语言 → Decoder → Cross-Attention → 生成下一个 Token

Decoder-only 模型没有独立 Encoder,也就没有原始意义上的 Encoder-Decoder Cross-Attention;不要把所有 Attention 都称为“Self-Attention”。


十一、训练与生成推理的数据流差异

11.1 训练:Teacher Forcing 与并行计算

训练一个长度为 nn 的序列时,可以一次输入:

x0 x1 x2 x3

并同时预测:

x1 x2 x3 x4

Causal Mask 保证第 ii 个位置不能看到未来,但 GPU 仍可以并行计算全部位置。

训练显存主要来自:

  • 参数;
  • 优化器状态;
  • 梯度;
  • 中间激活;
  • 注意力矩阵或其等价计算状态。

训练通常不使用生成式 KV Cache。因为训练需要对所有位置保留反向传播所需的计算图,且同一序列的所有位置一次计算更高效。把推理缓存机制直接套到训练上,可能造成额外内存和复杂的梯度依赖。

11.2 推理:逐 Token 自回归

生成过程如下:

sequenceDiagram
    participant C as 客户端
    participant M as 模型服务
    participant K as KV Cache
    C->>M: prompt
    M->>M: Prefill:处理完整 prompt
    M->>K: 保存各层 K、V
    M-->>C: 第一个生成 Token
    loop 每个后续 Token
        C->>M: 携带或复用会话状态
        M->>M: 只计算当前 Token 的 Q、K、V
        M->>K: 追加当前 K、V
        K-->>M: 返回历史 K、V
        M-->>C: 下一个 Token
    end

实际服务端通常不会让客户端传输 KV Cache,而是由服务进程按请求保存。图中的“客户端携带或复用会话状态”表示请求需要关联到服务端的会话状态,不代表缓存一定跨网络传输。


十二、KV Cache:缓存的到底是什么

12.1 没有缓存时的重复计算

假设当前已经有:

A B C

要生成 D,模型需要计算三个位置的 Q,K,VQ,K,V,最后使用 C 位置的 Query 读取 A、B、C

生成 E 时,输入变成:

A B C D

如果不使用缓存,模型会再次计算 A、B、C 的 Query、Key、Value。对于更长的输出,这会不断重复历史计算。

12.2 有缓存时的计算

在生成 D 时,保存:

KA:B:C,VA:B:CK_{A:B:C},\quad V_{A:B:C}

生成 E 时,只对新 Token D 计算:

qD, kD, vDq_D,\ k_D,\ v_D

并追加:

Kcache[Kcache;kD]K_{\text{cache}}\leftarrow [K_{\text{cache}};k_D]

Vcache[Vcache;vD]V_{\text{cache}}\leftarrow [V_{\text{cache}};v_D]

当前 Query 只需要计算:

oD=softmax(qDKalldk)Vallo_D= \operatorname{softmax} \left( \frac{q_DK_{\text{all}}^\top}{\sqrt{d_k}} \right)V_{\text{all}}

其中:

Kall=[KA,KB,KC,KD]K_{\text{all}}=[K_A,K_B,K_C,K_D]

Vall=[VA,VB,VC,VD]V_{\text{all}}=[V_A,V_B,V_C,V_D]

关键点是:

当前步仍然要让新的 Query 读取全部历史 Key 和 Value;KV Cache 省掉的是历史 K/V 的重新计算,不是消除当前 Query 对历史的注意力读取。

12.3 Cache 的形状

对一个请求、一个层、标准 MHA,常见形状为:

K,VRB×h×S×dkK,V\in\mathbb{R}^{B\times h\times S\times d_k}

其中:

  • BB:Batch 或请求批次;
  • hh:KV 头数;
  • SS:当前缓存序列长度;
  • dkd_k:每个头的维度。

在 MHA 中,KV 头数通常等于 Query 头数;在 GQA 或 MQA 中,KV 头数更少。

所有层都需要自己的缓存,因此总缓存元素量近似为:

2×L×B×S×hkv×dk2\times L\times B\times S\times h_{\text{kv}}\times d_k

其中:

  • LL:层数;
  • 前面的 2:分别表示 K 和 V;
  • hkvh_{\text{kv}}:Key/Value 头数。

若使用每元素 bb 字节的数据类型,缓存字节数约为:

2LBShkvdkb2LBS h_{\text{kv}}d_k b

例如:

层数 L       = 32
Batch B      = 1
序列长度 S   = 4096
KV 头数      = 8
每头维度     = 128
数据类型     = FP16,每元素 2 字节

则:

2×32×1×4096×8×128×22\times32\times1\times4096\times8\times128\times2

约为 512 MiB。这里没有计算缓存元数据、对齐、临时工作区和框架额外开销,因此实际占用可能更高。

12.4 Prefill 与 Decode

推理一般分为两个阶段:

Prefill

  • 一次处理完整 Prompt;
  • 计算所有 Prompt Token 的隐藏状态;
  • 生成并保存所有历史 K/V;
  • 通常计算量大,但可以并行。

Decode

  • 一次处理一个新 Token;
  • 只计算当前 Token 的隐藏状态;
  • 读取已有 KV Cache;
  • 将新的 K/V 追加到缓存;
  • 通常受显存带宽和逐步调度影响。

因此“输入很长但只生成一个 Token”和“输入很短但生成很多 Token”瓶颈不同:

  • 长输入主要影响 Prefill;
  • 长输出主要影响 Decode 和 KV Cache 增长。

十三、KV Cache 的状态管理和故障路径

KV Cache 不是模型参数,而是请求相关的运行时状态。一个缓存条目至少需要关联:

request_id
model_version
adapter_version(如果使用 LoRA 等适配器)
token 序列或位置状态
当前长度
KV 数据块
停止状态

如果把不同模型版本的缓存混用,即使张量形状完全相同,结果也可能错误,因为不同模型的投影矩阵、位置编码和层结构可能不同。

常见故障包括:

13.1 位置编号错误

如果缓存中已有 100 个 Token,却把新 Token 当作位置 0 重新应用 RoPE,模型看到的相对位置会错误,表现可能是:

  • 输出突然重复;
  • 语法结构断裂;
  • 长上下文质量明显下降;
  • 生成结果与无缓存版本不一致。

13.2 Cache 与 Token 序列不同步

如果服务端已经把某个 Token 的 K/V 追加到缓存,但客户端重试请求时又重复发送该 Token,就会出现:

缓存:A B C
重试输入:A B C C

这不是普通的网络重试问题,而是状态幂等性问题。生成服务需要明确请求是:

  • 继续已有前缀;
  • 替换会话;
  • 从某个缓存长度恢复;
  • 还是完整重算。

13.3 Batch 中序列长度不同

动态 Batch 中,每个请求的缓存长度可能不同。实现通常使用:

  • padding;
  • paged attention;
  • block table;
  • 变长索引;
  • prefix cache。

如果把不同请求的 KV 直接拼接而没有正确的长度和索引屏蔽,一个请求可能读取另一个请求的历史内容。这不仅是准确性错误,也是严重的数据隔离问题。

13.4 Beam Search 或分支生成

当多个候选序列从同一前缀分叉时,可以共享前缀 KV,但分叉后每个分支要维护自己的新增 KV。若原地修改共享缓存,可能造成不同候选相互污染。

13.5 请求取消和异常恢复

用户断开连接、超时或达到 Token 预算后,缓存必须释放或回收到缓存池。只在应用层标记请求结束而不释放底层块,会造成缓存泄漏,最终表现为显存持续增长和新请求排队。


十四、一个最小可运行的 Attention 与 KV Cache 示例

下面的代码使用 PyTorch 展示单头因果注意力。它不包含完整 Transformer,也不追求高性能,但能明确看到:

  1. 无缓存时对整个序列计算;
  2. 有缓存时只处理新 Token;
  3. 两种方式的结果应一致。
import math
import torch


def causal_attention_full(x, wq, wk, wv):
    """
    x:  [B, T, D]
    wq: [D, H]
    wk: [D, H]
    wv: [D, H]
    return:
        y: [B, T, H]
    """
    q = x @ wq
    k = x @ wk
    v = x @ wv

    scores = q @ k.transpose(-1, -2) / math.sqrt(q.size(-1))
    t = x.size(1)

    # 上三角为未来位置,禁止读取
    future = torch.triu(
        torch.ones(t, t, dtype=torch.bool, device=x.device),
        diagonal=1,
    )
    scores = scores.masked_fill(future, float("-inf"))

    weights = torch.softmax(scores, dim=-1)
    return weights @ v


def causal_attention_decode(x_new, wq, wk, wv, k_cache=None, v_cache=None):
    """
    x_new:   [B, T_new, D]
    k_cache: [B, T_old, H] or None
    v_cache: [B, T_old, H] or None

    这里为了展示机制,仍然计算 x_new 内部的注意力。
    自回归逐 Token 解码时通常 T_new=1,此时不需要额外的局部 causal mask。
    """
    q_new = x_new @ wq
    k_new = x_new @ wk
    v_new = x_new @ wv

    if k_cache is None:
        k_all = k_new
        v_all = v_new
    else:
        k_all = torch.cat([k_cache, k_new], dim=1)
        v_all = torch.cat([v_cache, v_new], dim=1)

    scores = q_new @ k_all.transpose(-1, -2) / math.sqrt(q_new.size(-1))

    # 当 T_new > 1 时,新片段的第 i 个位置不能读取片段内更晚的位置。
    old_t = 0 if k_cache is None else k_cache.size(1)
    new_t = x_new.size(1)

    if new_t > 1:
        future = torch.triu(
            torch.ones(new_t, new_t, dtype=torch.bool, device=x_new.device),
            diagonal=1,
        )
        prefix = torch.zeros(
            new_t, old_t, dtype=torch.bool, device=x_new.device
        )
        mask = torch.cat([prefix, future], dim=1)
        scores = scores.masked_fill(mask.unsqueeze(0), float("-inf"))

    weights = torch.softmax(scores, dim=-1)
    y_new = weights @ v_all

    return y_new, k_all, v_all


if __name__ == "__main__":
    torch.manual_seed(0)

    b, t, d, h = 1, 4, 8, 8
    x = torch.randn(b, t, d)
    wq = torch.randn(d, h)
    wk = torch.randn(d, h)
    wv = torch.randn(d, h)

    # 方式一:一次性处理完整序列
    y_full = causal_attention_full(x, wq, wk, wv)

    # 方式二:先处理第一个 Token,再逐个追加
    cache_k = None
    cache_v = None
    outputs = []

    for i in range(t):
        y_i, cache_k, cache_v = causal_attention_decode(
            x[:, i:i + 1],
            wq, wk, wv,
            cache_k, cache_v,
        )
        outputs.append(y_i)

    y_cached = torch.cat(outputs, dim=1)

    print("最大绝对误差:", (y_full - y_cached).abs().max().item())
    print("KV Cache 形状:", cache_k.shape, cache_v.shape)

在浮点误差范围内,最大绝对误差 应接近 0,且最后的 Cache 形状为:

torch.Size([1, 4, 8]) torch.Size([1, 4, 8])

每一步的状态变化是:

第 1 步:cache = K[x0], V[x0]
第 2 步:cache = K[x0, x1], V[x0, x1]
第 3 步:cache = K[x0, x1, x2], V[x0, x1, x2]
第 4 步:cache = K[x0, x1, x2, x3], V[x0, x1, x2, x3]

这个示例没有包含:

  • 多头拆分;
  • RoPE;
  • LayerNorm 或 RMSNorm;
  • FFN;
  • 输出投影;
  • 高性能 fused kernel;
  • paged KV Cache;
  • 量化缓存。

因此它用于验证数学关系,不应直接作为生产推理内核。生产框架通常会使用 fused attention、连续批处理、分页缓存和专用 CUDA Kernel。


十五、常见误解与诊断方法

15.1 “Attention 就是查数据库”

Attention 的确具有“按相关性读取信息”的行为,但 Key、Value 不是外部数据库记录,而是当前网络层根据输入动态计算出的连续向量。它没有天然的精确检索保证,也不自动执行符号匹配。

如果模型需要可靠地读取外部事实,通常需要检索系统、工具调用或数据库查询;不能仅凭 Attention 的存在推断事实一定正确。

15.2 “KV Cache 会减少所有推理计算”

KV Cache 只减少历史 Token 的 K/V 投影和相关中间计算。当前 Query 仍然要与历史 Key 做点积,并对历史 Value 做加权求和。因此随着上下文长度增加,单步 Decode 的注意力读取成本仍会增长。

另外,FFN、当前 Token 的 Query 投影、输出投影和采样仍然需要执行。

15.3 “上下文长度等于 KV Cache 容量”

上下文长度是模型允许参与计算的 Token 范围;KV Cache 是某个请求当前保存的运行时张量。一个服务可以:

  • 只保存滑动窗口;
  • 截断旧上下文;
  • 使用分层或分页缓存;
  • 为某些层采用特殊注意力结构。

因此“模型上下文窗口 128K”并不意味着每个请求都必须永久保存 128K Token 的完整 FP16 KV。

15.4 “加大 Temperature 就让模型更聪明”

Temperature 只改变采样分布。它不会增加模型参数、上下文信息或事实知识。Temperature 过高常表现为:

  • 逻辑跳跃;
  • 格式不稳定;
  • 罕见 Token 增多;
  • 结构化输出解析失败。

Temperature 过低则可能增加重复和模板化。应结合任务目标和评测结果调整,而不是把它当作能力开关。

15.5 “LayerNorm 是对整个句子做归一化”

标准 LayerNorm 通常对每个 Token 的隐藏维度独立计算均值和方差。它不是把一个句子的所有 Token 混成一个统计量,也不是 BatchNorm。误解归一化轴会直接导致实现结果不同。

15.6 “训练损失低就说明生成质量高”

交叉熵衡量的是训练数据分布上的下一 Token 概率。低损失不保证:

  • 指令遵循;
  • 事实正确;
  • 安全策略;
  • 长文本一致性;
  • 结构化输出可解析;
  • 未见任务泛化。

完整生产评测应同时覆盖离线任务指标、生成质量、拒答策略、权限边界、成本和延迟。


十六、生产系统中的模型、数据、权限与成本

Transformer 推理不是孤立的数学函数,而是生产系统的一部分。

16.1 模型和版本

以下内容都会影响输出,不能只记录一个模型名称:

基础模型版本
Tokenizer 版本
Chat Template
LoRA / Adapter 版本
量化配置
推理框架版本
采样参数
安全策略版本

Tokenizer 变化可能导致相同文本映射到不同 Token 序列,进而改变位置、KV Cache 和模型输出。模型升级后,旧的 Prefix Cache 或 KV Cache 不应默认复用,除非系统明确验证了兼容条件。

16.2 数据和评测

预训练、指令微调和对齐阶段优化的目标不同:

  • 预训练学习语言和世界知识分布;
  • 指令微调学习任务格式和人类指令;
  • 对齐训练使输出更符合偏好、安全和行为约束。

评测时应固定:

  • Tokenizer;
  • Prompt 模板;
  • 最大输入与输出长度;
  • 解码策略;
  • 随机种子或采样配置;
  • 失败重试规则。

否则不同版本之间的结果差异可能来自推理配置,而不是模型本身。

16.3 权限和数据隔离

KV Cache、Prompt Cache、日志和中间输出都可能包含敏感信息。多租户服务必须保证:

  • 不同租户不能共享未验证的前缀缓存;
  • 请求取消后缓存和日志按策略清理;
  • 工具调用权限与模型生成权限分离;
  • 模型不能因为上下文中出现某段文本就自动获得数据库或文件系统权限。

Attention 可以读取上下文,不等于模型拥有外部系统的访问权。权限必须由服务端工具层、身份认证和策略引擎强制执行。

16.4 成本与容量

推理成本通常同时受以下因素影响:

  • 输入 Token 数;
  • 输出 Token 数;
  • Batch 大小;
  • 模型参数量;
  • KV Cache 长度;
  • 数据类型;
  • 并发量;
  • Prefill 与 Decode 的比例。

KV Cache 通过减少重复计算提高吞吐,但会增加每个活跃请求的显存占用。若缓存池被长请求占满,新请求可能出现排队、拒绝或被迫卸载。系统需要明确超时、最大上下文、最大输出和缓存回收策略,而不能只追求单请求速度。


十七、把整个 Transformer Block 串起来

以 Pre-LN、decoder-only、带因果掩码的一个 Block 为例,输入为:

X(l1)RB×T×dmodelX^{(l-1)}\in\mathbb{R}^{B\times T\times d_{\text{model}}}

先做第一层归一化:

U=Norm1(X(l1))U=\operatorname{Norm}_1(X^{(l-1)})

计算多头因果自注意力:

A=MHAcausal(U)A=\operatorname{MHA}_{\text{causal}}(U)

第一次残差相加:

Y=X(l1)+AY=X^{(l-1)}+A

再归一化:

Z=Norm2(Y)Z=\operatorname{Norm}_2(Y)

逐 Token 计算 FFN:

F=FFN(Z)F=\operatorname{FFN}(Z)

第二次残差相加:

X(l)=Y+FX^{(l)}=Y+F

重复 LL 次后:

H=X(L)H=X^{(L)}

对每个位置的隐藏状态映射到词表:

Logits=HWlm+b\text{Logits}=HW_{\text{lm}}+b

训练时对所有位置计算交叉熵;生成时通常只取最后一个位置的 logits,选择下一个 Token,再进入下一轮。

这一过程中的职责分工是:

Embedding       提供 Token 的初始语义表示
位置表示         提供顺序和相对位置信息
Attention        在位置之间交换信息
Causal Mask      阻止读取未来
残差连接         保留并累积已有表示
归一化           控制隐藏状态尺度
FFN              对每个位置进行非线性通道变换
LM Head          将隐藏状态映射到词表分数
解码策略         将概率分布转成离散 Token
KV Cache         避免生成时重复计算历史 K/V

Transformer 的核心并不是某一个单独公式,而是这些组件在训练和自回归推理中的组合。Attention 决定当前位置如何读取上下文;残差和归一化决定深层表示如何稳定演化;Decoder 和 Causal Mask 决定模型如何遵守下一 Token 预测约束;解码器把概率转成文本;KV Cache 则把同一模型从“每一步重算完整前缀”变成“缓存历史、只计算增量”的推理系统。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。