Transformer核心公式深度拆解:从Attention机制到KV cache工程实践
大模型的军备竞赛已经持续了好几年。百模大战也好开源闭源之争也好无论模型叫什么名字、用的是什么训练策略它们在架构层面几乎都指向同一个来历2017 年那篇不到 12 页的论文带来的 Transformer 结构。后来各种大模型能够统一在“预训练 微调 对齐”这条范式下本质上也是因为 Transformer 的编码器、解码器成了事实标准。这篇文章想聊的不是什么新框架也不是某个大模型的评测结果而是 Transformer 论文里那条最核心的公式[ Attention(Q, K, V) softmax(\frac{QK^T}{\sqrt{d_k}})V ]全球大模型都在用这条公式它决定了 token 之间如何交互、信息如何被抽取、上下文如何被建模。但奇怪的是很多工程同学在模型部署和微调时能熟练操作各类工具真被问到这条公式里为什么要除以根号 d_k或者 causal mask 在计算图上到底怎么作用时却容易卡壳。如果你正处在这样的阶段那这篇文章可以帮你把这些概念串起来。1. 这篇文章真正要解决的问题先说一句直接判断Transformer 这条公式不只是一段数学表达式它同时是模型结构、训练效率和推理成本的分水岭。为什么这么说因为大模型的所有关键行为都由它衍生出来token 之间的相关性怎么算是这条式子决定的为什么长文本会带来推理显存压力也是因为这条式子里的 K、V 需要缓存为什么算力都被“吃掉”在矩阵乘法上还是源于这条式子为什么模型可以并行处理整段序列而 RNN 不行依然要回到这条式子的整体结构。很多人的困惑在于这条公式看起来并不复杂就是一个 softmax 套在点积外面很简短但围绕它的衍生工程概念却很多多头注意力、因果掩码、温度系数、FlashAttention、KV cache、RoPE这些概念如果孤立地学会非常散很难融会贯通。但如果以这条公式为锚点很多问题都能串起来理解。这篇文章适合这几类读者正在学习大模型原理、准备入门训练或微调的同学。已经会使用模型部署工具但想知道“模型在 GPU 上到底算了什么”的工程师。需要排查大模型推理显存溢出、上下文窗口受限等问题的开发者。想系统理解 QKV、注意力掩码、KV cache 等概念的技术人。读完这篇文章你可以做到理解这条公式中 Q、K、V 各自的含义知道为什么要除以 (\sqrt{d_k})能解释 softmax 在这里起到的作用理解为什么 GPT 这类模型只能从左向右看知道推理时的 KV cache 到底缓存了什么。这篇文章不展开讲全部训练细节也不会堆叠复杂数学推导而是按“公式拆解 — 结构演进 — 工程落地 — 问题排查”这条思路来讲。2. 这个公式为什么能统治大模型从 RNN 时代过来的老开发者应该对序列建模的痛苦记忆犹新。过去处理文本主流方案是 LSTM、GRU 这类循环神经网络。它们的计算特点是当前时刻的隐藏状态依赖上一个时刻的输出天然只能一步一步算。这种串行结构带来的问题很明显序列一长梯度容易消失早期信息就传不过来训练无法充分并行GPU 利用率上不去长距离依赖的建模能力有限。RNN 之所以慢不只是结构上的串行更关键的是它把“记忆”压缩在一个固定维度的状态向量里。文本长得离谱时这个状态向量根本存不下所有信息。你让一个固定大小的容器装下一百年前的信息它很快就会不堪重负。Transformer 的做法是换了一种完全不同的思路不再用一个状态向量逐步传递信息而是让序列中的每一个 token 直接和所有其他 token 计算相关性。这就相当于每个词都可以同时“看到”全句的每个词信息不再需要逐级传递而是一步到位。Q、K、V 这套概念其实借用了信息检索里的思想。你可以把它理解成一个词库查询过程QQuery查询向量你脑子里想找什么。KKey键向量每个词对外展示的“标签”。VValue值向量每个词真正承载的内容。一个直观的类比是文件检索你带着一个查询词先和每个文件的标签做匹配计算出相似度再按相似度把最相关文件的内容抽取出来。在自注意力机制里每个 token 同时扮演了三重角色作为询问者去查询其他 token作为被查询者提供自己的标签作为内容源提供实际的信息。这种机制使文本中的每个词都能根据上下文动态调整自己的表征同一个词在不同语境中自然会有不同的向量表示。这条公式是整个大模型架构的基石也直接带来了一个副产品上下文容量变成了影响模型效果和成本的重要指标。绝大多数大模型都采用这种架构它绕过了循环结构用自注意力和位置编码完成序列建模。在 2017 年那个时间点这条公式最令人震撼的地方在于它证明了一件事序列建模最强的路径不见得是设计更复杂的“记忆单元”而是把全序列的信息放在一个开放交互的空间里让模型自己学出该关注什么。这种“简单且规模友好”的设计是它后来成为大模型共识的重要原因。3. 逐项拆解 Attention 公式理解了背景下面把公式拆开来看。真实的计算过程没有想象中那么玄妙宏观上就是三次矩阵乘法加一次缩放、一次归一化。3.1 Q、K、V 是怎么来的首先需要明确一点Q、K、V 不是原始输入原始输入 X 必须先经过三层线性变换才能变成 Q、K、V[ Q XW_Q,\quad K XW_K,\quad V XW_V ]这里的 (W_Q)、(W_K)、(W_V) 都是可训练的权重矩阵分别投影到各自的向量空间。它们的作用是让模型能学习到“查询、匹配、内容”三种不同的特征空间。输入 X 的维度通常是 [batch_size, seq_len, hidden_size]。经过这三个权重矩阵后Q、K、V 的形状不变只是每个 token 的向量表达被映射到了不同空间。为什么要分开三个空间因为查询相关性和内容提取不应该共用同一个向量。如果把一个向量既当查询又当内容模型容易把“我想找谁”和“我实际携带什么”混淆。分开投影后模型可以分别学习这两类信息表达更灵活。这是这条公式一个重要但容易忽视的设计。3.2 点积 QK^T 是在算什么Q 和 K 做点积得到的分数代表“匹配程度”。Q 的第 i 行代表第 i 个 token 的查询向量K^T 的第 j 列代表第 j 个 token 的键向量两者做点积得到的就是第 i 个 token 对第 j 个 token 的注意力分数。如果两个向量在空间中方向越接近、长度越匹配点积越大表示相关性越强。这个分数矩阵的形状是 [seq_len, seq_len]就是常说的注意力矩阵。它的每一行代表一个 token对所有其他 token 分配了多少注意力。这个矩阵是理解一切注意力机制行为的基础它直接决定了信息如何流动。下面用一个最小示例来看。在实际代码中经常见到的是import numpy as np def attention_scores(Q, K): # Q, K shape: (seq_len, dim) return np.dot(Q, K.T)如果 Q 和 K 的维度是 3计算两个 token 之间的点积q [1.0, 0.0, 1.0] k1 [0.8, 0.2, 0.8] k2 [0.0, 1.0, 0.0] score1 sum(a * b for a, b in zip(q, k1)) # 1.6 score2 sum(a * b for a, b in zip(q, k2)) # 0.0score1 大于 score2说明第一个 token 和 k1 更相关。但这里的问题也随之而来如果向量维度很大得到的分数区间会很大softmax 几乎会把所有概率都压到最大值那个位置上。所以必须进行缩放。3.3 为什么除以根号 d_k这是公式里看似简单、实际非常关键的一步。假如 Q 和 K 的向量维度是 (d_k)且向量中的每个元素独立服从均值 0、方差 1 的分布那么两个向量的点积[ Q \cdot K \sum_{i1}^{d_k} q_i k_i ]这个结果的均值为 0方差却是 (d_k)。也就是说维度越高点积结果的方差越大。方差大意味着某些分数会非常大而另一些会非常小。计算 softmax 时如果某个分数很大softmax 的输出会变得非常接近 one-hot 分布最大的那个元素几乎拿到全部概率其他位置趋于 0梯度也会变得极小。模型一旦进入这种状态学习几乎停摆因为反向传播时那些接近 0 的位置很难把梯度传回去。除以 (\sqrt{d_k}) 的目的就是缩放点积使方差回到一个稳定的范围。如果不做这个缩放大维度场景下注意力分布容易退化训练会出现不稳定的情况。这是这条公式里最值得深挖的细节之一。很多讲解都把这一步当成约定俗成但实际上它决定了训练稳定性。从业界经验看这个缩放因子可以视为一个固定温度系数让点积结果保持在概率计算更舒服的区间内。3.4 softmax 在这里做了什么事情softmax 的作用是在缩放后的分数矩阵上按行做归一化让每一行所有分数变成非负且和为 1 的概率分布。本质上它把“哪些 token 值得关注”表达成一组权重。[ softmax(z)_i \frac{e^{z_i}}{\sum_j e^{z_j}} ]softmax 有一个重要特性它的输出总是正的即使原始分数是负数经过指数运算后也会变成正数。这意味着模型始终会给每个 token 分配一个微小的注意力而不是完全忽略某个 token。从这个角度看注意力机制更像是一种“软性”的信息筛选它不会彻底丢弃信息而是把信息按重要程度重新组合。模型训练的过程同时也是在学这一组权重。需要特别提醒的是 softmax 在数值计算上的一个坑如果输入分数里有很大的正值比如 50那么 e 的 50 次方会非常大在 float16 下很容易溢出。主流实现一般会先减去每行的最大值再做指数运算def softmax(x): x_max np.max(x, axis-1, keepdimsTrue) exp_x np.exp(x - x_max) return exp_x / np.sum(exp_x, axis-1, keepdimsTrue)3.5 最后乘 V 是在抽取信息得到注意力权重后下一步是乘 V。这实际是一个加权求和的过程每个 token 的新表示是所有 token 的 V 向量按注意力权重混合后的结果。[ Attention(Q,K,V) P \cdot V ]假设第 i 个 token 的注意力分数经过 softmax 后得到权重 0.7、0.2、0.1那么第 i 个 token 的新表示就是 0.7 倍的 V[0]加上 0.2 倍的 V[1]再加上 0.1 倍的 V[2]。这个加权求和非常好地体现了局部与整体的统一注意力权重关注的是“谁更重要”V 向量提供“要传递的具体信息”。两者相乘才能完成真正的上下文信息融合。到这里注意力的整体计算就完整了。用简单的 Python 实现可以把所有步骤写在一起import numpy as np def softmax(x): x_max np.max(x, axis-1, keepdimsTrue) exp_x np.exp(x - x_max) return exp_x / np.sum(exp_x, axis-1, keepdimsTrue) def attention(Q, K, V, maskNone): d_k Q.shape[-1] scores np.matmul(Q, K.transpose(0, 1, 3, 2)) / np.sqrt(d_k) if mask is not None: scores scores mask weights softmax(scores) return np.matmul(weights, V)在真实框架中一般会直接用 PyTorch 或 TensorFlow 的实现但核心流程完全一致。4. 从 self-attention 到多头注意力基础公式解决了 token 之间的交互问题但只用一组 Q、K、V 去建模复杂的语言表达能力仍然不足。语言中的关系是多种多样的相邻词之间的语法关系、远距离依赖的指代关系、语义上的上下位关系这些不同层面的信息往往需要不同的表示空间。多头注意力Multi-Head Attention的解决方案是不使用一组注意力而是使用多组并行的注意力每组称为一个头。每个头有自己独立的 Q、K、V 投影矩阵因此可以在不同表示子空间里捕捉不同的关联模式。举个例子一个头的注意力可能集中在句法邻近的词上另一个头可能集中在长距离指代关系上。多个头的结果拼接在一起再经过一层线性变换得到最终输出。多头注意力的宏观公式是[ MultiHead(Q,K,V) Concat(head_1,...,head_h)W_O ]其中每个 head 都是一次独立的注意力计算。放到完整公式里其实是把原来那一条公式平行复制了多次。这个设计不是因为单次注意力不够强而是为了让模型拥有多种“观察角度”。从计算效率上看多头注意力增加的计算量并不是随头数线性增加。因为每个头的维度会减少比如隐藏层是 76812 个头时每个头只有 64 维整体计算量基本与单头保持在同一量级。这样就在不大幅增加成本的前提下提升了模型表达能力。需要考虑的一点是在推理阶段如果 K、V 需要缓存缓存的也是所有头的 K、V。因此多头数量直接影响 KV cache 的大小。5. 位置编码与注意力公式的互补关系看到这里有人可能会产生疑问这条公式本身并没有显式建模 token 的先后顺序。把句子的任意两个 token 调换位置它们两两之间照样能算注意力结果都不受影响。这说明纯注意力机制对语言顺序是完全“无感”的。但语言是强顺序相关的。“张三打了李四”和“李四打了张三”词完全一样意思却相反。所以必须在架构里补充位置信息。这就是位置编码Positional Encoding的由来。最著名的做法是 Transformer 论文里提出的三角函数位置编码[ PE(pos, 2i) \sin(\frac{pos}{10000^{2i/d}}) ] [ PE(pos, 2i1) \cos(\frac{pos}{10000^{2i/d}}) ]它将绝对位置信息通过正弦余弦函数叠加到词向量上。后来的大模型则更倾向使用 RoPE旋转位置编码它把位置信息编码成旋转矩阵并在 Q、K 做点积时自然地把相对位置信息融入注意力分数。RoPE 之所以流行一个重要原因是它对相对距离的表达更自然对推理时的长度外推也更友好。这里想强调一个关键认知位置编码并不是 Attention 公式内部的组成部分但它通过修改 Q、K 的输入间接影响了注意力分数的计算结果。工程上调整上下文窗口、做长文本推理时很多问题既可能出自训练数据也可能出自位置编码的外推能力。理解这一点有助于避免把所有问题都归因于模型参数量。当前大模型大多采用类似 GPT 的因果解码器结构每个 token 只能关注自己左边的 token不能看右边。这意味着我们计算公式中的注意力时注意力矩阵是一个下三角形状的有效区域。在自回归生成场景下信息只能从过去流向未来这样才能保证模型预测下一个 token 时不会“偷看”答案。在工程实现中通过让每个 token 与它后面 token 的注意力分数加上一个负无穷来实现。6. 公式之外的工程现实部署与推理很多做模型部署的同学觉得论文公式离实际很远其实恰好相反。这条公式直接决定了大模型推理时的算力和显存特点也解释了为什么“大模型部署”一直是热门话题。它并不只是模型文件拷贝到 GPU 上那么简单显存占用、解码速度、可支持的上下文长度本质上都受到注意力计算的制约。6.1 全套大模型推理流程现在主流的生成流程可以概括为两个阶段预填充阶段Prefill用户输入的一整段提示词与 KV cache 一起做一次完整的前向计算生成第一个 token。解码阶段Decode每生成一个新 token需要将其与之前的 token 一起更新 KV cache进行下一次前向计算。在解码阶段由于每个新的 token 只能看到之前的 token传统注意力计算每次都要对整个历史序列执行一次完整的 softmax。这样会带来明显问题序列越长计算量和显存占用会越来越大。6.2 KV cache 缓存内容重要的问题是每次生成新 token 时注意力计算其实并不需要所有历史信息只需要新 token 的 Q以及所有历史 token 的 K 和 V。也就是说K 和 V 一旦计算完成后续采样过程中可以复用不需要每步重算。这就是 KV cache 的基本思想。KV cache 增大显存会被大量消耗。以 GPT 这类模型为例KV cache 的显存占用大约正比于[ 2 \times batch_size \times seq_len \times num_layers \times num_heads \times head_dim \times precision_bytes ]其中 2 代表 K 和 V 两份缓存。可以看到序列越长KV cache 占用的显存就越高。这也是长文本生成时显存不足的一个非常重要但容易被忽略的原因。6.3 为什么 vLLM、FlashAttention 这些工作如此重要当 KV cache 很大时传统注意力实现的显存开销会很大瓶颈尤其明显。业界因此出现了一系列优化方向FlashAttention通过分块计算注意力避免把完整的注意力矩阵写入显存从而减少显存占用并提升计算效率。vLLM 这类推理框架采用 PagedAttention 思想将 KV cache 划分为固定大小的块进行管理类似操作系统中的分页机制减少显存碎片和浪费。推测解码Speculative Decoding用小模型先草拟多个 token再用大模型一次验证降低延迟。这些优化本质上是在解决同一个问题的不同侧面Attention 公式本身虽然简明但当输入序列变长、并发请求变多它的计算和存储需求会急剧增长。这也是为什么“部署一个大模型”绝不等于“能启动就行”还要考虑吞吐、延迟和可支持的上下文长度。7. 完整代码实现与效果验证理论部分讲完后实践环节必不可少。下面用最小 Python 代码实现一个简化版的自注意力模块方便自己调试和理解。这里不使用大型框架而是用 NumPy 实现以便观察每一步输出。完整示例import numpy as np def softmax(x): x_max np.max(x, axis-1, keepdimsTrue) exp_x np.exp(x - x_max) return exp_x / np.sum(exp_x, axis-1, keepdimsTrue) class SimpleSelfAttention: def __init__(self, d_model, d_k, d_v): self.d_k d_k self.d_v d_v self.W_Q np.random.randn(d_model, d_k) * 0.1 self.W_K np.random.randn(d_model, d_k) * 0.1 self.W_V np.random.randn(d_model, d_v) * 0.1 def forward(self, X): # X shape: (seq_len, d_model) Q np.dot(X, self.W_Q) K np.dot(X, self.W_K) V np.dot(X, self.W_V) scores np.dot(Q, K.T) / np.sqrt(self.d_k) weights softmax(scores) output np.dot(weights, V) return output, weights np.random.seed(42) X np.array([ [1.0, 0.0, 0.5, 0.2], [0.0, 1.0, 0.3, 0.1], [0.5, 0.5, 1.0, 0.8] ]) attention_layer SimpleSelfAttention(d_model4, d_k4, d_v4) output, weights attention_layer.forward(X) print(Attention Weights:) print(weights) print(Output:) print(output)运行后输出结果是一个 [seq_len, d_v] 的矩阵同时可以看到注意力权重矩阵。简单封装后可以把代码放进simple_attention.py中执行python simple_attention.py预期输出中会看到一个 3x3 的注意力权重矩阵每一行的数值之和接近 1体现 softmax 的行归一化效果。输出矩阵的每个向量都是对输入序列中各个 token 信息按注意力权重做加权平均后的结果。如果你把 Q、K 换成线性层再把多头循环加进来就离 PyTorch 的nn.MultiheadAttention不远了。真正的框架实现还会加入头维度的 splitting 逻辑以及梯度计算但基础结构是一致的。做这个小实验时如果分数维度设置得大一些例如 d_k128即使不除 (\sqrt{d_k})也能观察到 softmax 概率分布趋向于尖锐最大值接近 1其他位置接近 0。这比直接用文字描述更有说服力。8. 常见问题与排查方法在实际学习和项目中围绕 Attention 公式的疑问很多下面以表格形式整理高频问题并给出定位思路。问题现象可能原因排查方式解决方案训练时 loss 波动剧烈或梯度异常Attention 分数过大softmax 区域饱和检查 QK^T 的数值范围观察是否出现极大值检查是否已做缩放必要时调整初始化方式使用 LayerNorm 稳定输入长文本推理时显存溢出KV cache 占用过大监控显存变化观察是否随生成长度线性增长使用 PagedAttention、KV cache 量化或减少并发 batch模型生成重复内容注意力分布过于集中或温度设置偏低检查生成参数与注意力权重观察高概率 token 是否集中在固定区域调整 temperature、top-p加入 repetition penalty训练时注意力分数出现 NaNsoftmax 输入过大在 float16 下溢出检查模型是否使用混合精度观察 loss 数值对注意力分数进行数值稳定化必要时梯度裁剪多头注意力效果与单头无明显差异或训练缓慢多头维度被压得过低或初始化不当查看各头注意力图是否相似调整 head 数量和 head_dim 的平衡检查学习率推理结果与训练时不一致推理长度超过训练长度位置编码外推能力弱对比不同长度输入的困惑度或输出使用支持长度外推的 RoPE 变体避免直接超长输入因果语言模型生成时漏看历史信息mask 未正确应用到注意力层检查 attention mask 的 shape 与取值在 scores 上加上上三角掩码矩阵掩码位置填入负无穷在工程排查中最基础也最容易被忽略的一点是先确认当前问题到底出在“模型结构”还是“推理实现”。遇到过同学将显存不足直接归因于模型文件太大其实很多时候是 KV cache 在高并发下被放大。这时先用降低 batch size 的方式验证再考虑换推理框架会更高效。9. 最佳实践与工程建议结合多年大模型应用与推理经验针对 Attention 公式延伸到实际项目总结几条建议。9.1 先理解结构再选择部署工具不要为了追新而换工具。很多项目只是需要 API 调用未必需要本地部署需要本地部署的项目如果只是单卡测试用轻量框架快速跑通可能更合适。部署前先估算模型参数量和显存需求显存需求量级大致包含模型权重、优化器状态训练时、激活值、KV cache。对于推理尤其要估算 KV cache 大小。一句话结论如果只是常规个人实验优先选择低比特量化与成熟推理框架如果是生产环境再考虑 vLLM、SGLang 等高性能方案。9.2 重视上下文长度与 KV cache 的关系大模型上下文很长是优势但长上下文的推理成本同样很高。生产环境中把上下文控制在合理范围内往往比一味拉长更能控制成本和延迟。一些项目组会把用户的历史会话摘要化而不是把完整原始消息全部塞进上下文这种工程手段在成本控制上非常重要。9.3 做微调时注意学习率和数值稳定性微调大模型时因为预训练权重已经收敛学习中更容易受到注意力分数数值扰动的影响。建议保持比较低的学习率配合梯度裁剪并在训练初期观察 loss 是否出现尖刺。不要因为看到开源代码里没有梯度裁剪就认为它不重要。9.4 多留意 mask 和序列边界的实现在实际业务场景中如果输入包含多个会话片段或者需要在同一个序列里塞多个样本padding 和 attention mask 处理稍有偏差模型就会跨边界混淆信息。建议把 mask 逻辑单独抽出来做单元测试确保被 mask 的位置不会参与注意力分数的 softmax 运算。9.5 建立从公式到性能指标的关联大模型的性能指标并不只有 loss。吃透公式之后你会知道模型结构固定时KV cache 与序列长度是线性关系如果瓶颈在算力重点看打点来优化矩阵运算效率如果瓶颈在显存重点看 KV cache 量化和剩余显存大小。把这些对应关系建立起来你再去看任何模型部署框架的代码或者官方文档时会轻松很多。10. 总结与后续学习方向这篇文章用一条核心公式把它们串在了一起。到现在已经可以看懂这条公式背后不少有价值的问题为什么除以根号 d_k为什么用 softmax为什么需要多头注意力为什么推理阶段要缓存 KV为什么超长文本会影响显存和延迟。这些都是大模型工程里真实会遇到的问题而不是停留在概念层面的名词。下一步建议按这个顺序继续深入手写一次 mini Transformer把 attention、多头、LayerNorm、前馈网络全部实现一遍。用 PyTorch 的nn.Transformer或nn.MultiheadAttention跑一个简单翻译或文本生成 Demo。阅读 FlashAttention 的论文了解对注意力计算过程的显存优化方式。在自己训练的小模型上加不同位置编码对比效果差异。接入一个推理框架观察不同 batch size、不同输入长度下的显存和耗时曲线。如果你正在做大模型部署或微调建议把这条公式抄在工位旁边的白板上。看懂了它你再看任何大模型相关代码都会多一层“原来这里在做这件事”的感觉。建议先收藏备用继续动手跑一个自己的最小实现。