资讯详情

注意力之外:从位置编码到KV Cache的LLM工程实战

📅 2026/10/12 3:42:30 | 华诺云谱 👁 阅读
注意力之外:从位置编码到KV Cache的LLM工程实战
写MiniMind学习笔记第四篇正卡在注意力机制刚学完的节点上。标题叫“注意力之外位置、记忆、省显存、省时间、深加工”你要是只看字面可能以为这是杂谈其实这六个词几乎覆盖了从Attention到完整LLM之间最容易被低估的六个工程与算法问题。我自己在实现MiniMind之前天真地以为把自注意力代码敲出来就能复现一个能跑的模型结果真正训练、推理、调显存的时候才发现99%的时间都花在这“六个之外”的东西上。这篇笔记不讲注意力本身而是把位置编码、KV Cache、显存优化、时间优化、FFN和归一化这些“配角”逐个拆开所有内容围绕一个几十M参数的教学级LLM展开适合自己也写过或准备写小型GPT模型的人阅读。1. 位置编码给注意力补上坐标1.1 为什么注意力天生“失忆”自注意力机制本质上是对一组向量做加权求和权重大小由各个位置之间的相似度决定。这里有个致命特性你把序列任意打乱顺序注意力输出的结果不会变因为点积和加权求和都是集合运算不关心元素顺序。模型看到“小明吃饭”和“吃饭小明”在没有额外信息的情况下会认为它们是同一个句子这就是所谓的置换不变性。语言是强有序的“我爱你”和“你爱我”意思完全相反。所以必须给每个token注入位置信息让注意力能区分位置。在Transformer时代最经典的做法是绝对正弦位置编码给每个位置pos分配一个向量其中第2i维和第2i1维分别是sin(pos / 10000^(2i/d))和cos(pos / 10000^(2i/d))。我当时学到这里有个疑惑为什么要用不同频率的正弦函数而不是直接用一个递增的整数后来想通了——整数位置编码的取值范围会随序列长度无限增加模型难以外推。而正弦函数的值域固定在[-1,1]且不同频率的组合能形成类似“二进制编码”的多尺度模式低频维度编码大范围位置高频维度编码细微位置差异。这让模型更容易学会相对位置关系比如“第i个token在第j个token前面”这种信息。1.2 从绝对位置到相对位置RoPE绝对位置编码有个问题它把位置当作独立的坐标直接加到token上但注意力更多时候关心的是词与词之间的相对距离。例如转述、指代、句法依赖本质上都依赖相对距离而不是绝对坐标。后来出现的旋转位置编码RoPE就在解决这个问题。RoPE的核心理念是将位置信息编码成旋转矩阵。对二维向量(x1, x2)位置m的旋转可以写成(x1, x2) (x1 * cos(mtheta) - x2 * sin(mtheta), x1 * sin(mtheta) x2 * cos(mtheta))这个变换的本质是对向量进行旋转。如果把两个token的查询和键都做对应的旋转那么它们的点积结果会自动包含“位置差m-n”的项q_m^T k_n (R(q_m))^T (R(k_n)) q^T k * cos((m-n)*theta) ...于是注意力分数只依赖相对位置差(m-n)模型在训练中更容易学到“相邻更近”等规律。而且RoPE是乘性作用在向量上不会改变原始语义向量长度外推能力也更好。MiniMind中我用的是LLaMA风格实现将Q和K的最后一维分成前后两半分别用一组旋转角度进行旋转。代码写起来不到二十行但有两个容易踩坑的点一是旋转角度要预热到float32否则在低精度下会累积数值误差二是通常只在Q和K上做RoPEV不做因为V只是被加权的原始信息不需要携带位置约束。1.3 位置编码选择与实战如果不想自己实现RoPE也可以用ALiBi这种更粗暴的方案不修改向量而是在注意力分数上直接减去一个“距离惩罚项”分数矩阵中位置i和j的得分减去m*|i-j|。ALiBi在推理时可以平滑外推但它在短序列上的表达力略弱于学习到的位置编码。我的建议是如果是新项目直接上RoPE它现在几乎成了开源LLM的事实标准。实测MiniMind在上下文长度128时绝对位置编码和RoPE差距不明显但把推理长度拉到训练时的两倍以上RoPE的优势会立刻体现出来。还有一个细节位置编码是加在embedding上还是加在Q/K上这决定了信息注入的位置。早期Transformer是加在embedding上现在主流是RoPE加在Q/K上注入位置给注意力“输入”更直接也避免把原始词向量污染成“位置语义”混合体。2. 记忆KV Cache用显存换时间2.1 自回归生成的重复计算问题训练完之后我们使用LLM时是在做自回归生成每预测一个新token把它拼到输入末尾再重新跑一遍整个序列计算所有位置的自注意力。这样做有个巨大的浪费前n个token已经算过一遍Q、K、V了生成第n1个token时又要重新算一遍。在Transformer解码阶段当前token的Q只能与前面所有token的K和V做注意力计算后面多出来的token不会改变前面token的历史。所以我们可以把每一层的历史K和V缓存下来下次生成时只算当前token的新K、V再拼接到缓存后面注意力分数只与缓存做一次矩阵乘法就行。这就是KV Cache又叫键值缓存。实现KV Cache非常简单为每层分配两个空的tensorshape是(batch, num_heads, max_seq_len, head_dim)或(batch, max_seq_len, num_heads, head_dim)推理时每步把新的K、V写进对应位置。但这里有一个新手常犯的错直接在整个max_seq_len上做注意力把未填充位置也算了进去。必须在注意力mask里把未写入的位置设为负无穷否则模型会看到一堆随机初始化的0向量。2.2 KV Cache的显存账本KV Cache不是免费的午餐它把“每步节省的计算”转换成了“持续增长的显存占用”。显存大小的公式是KV Cache字节数 2K和V × batch_size × num_layers × num_heads × max_seq_len × head_dim × bytes_per_element拿MiniMind的一个典型配置举例d_model512num_heads8head_dim64num_layers6生成序列长度为1024使用FP162字节2 × 1 × 6 × 8 × 1024 × 64 × 2 201326592 字节 ≈ 192 MB如果batch_size变成8就是1.5 GB。这个数字相当惊人因为MiniMind的模型权重本身可能只有200多MKV Cache在长上下文场景下会直接超过权重。更夸张的是在2048长度下批量并发推理显存可能最先被KV Cache撑爆而不是模型参数。所以在正式做推理服务时几乎没人不优化KV Cache。最直接的方法是缓存精度从FP16降到INT8牺牲一点质量换一半显存或者对KV做压缩只存部分历史信息。另一个常见工程技巧是PagedAttention把KV Cache分页管理类似操作系统虚拟内存按需分配减少碎片化但这是另一个话题了。2.3 MQA/GQA从多头到分组多头既然KV Cache主要被“头数”和“层数”放大一个聪明的思路是让多个查询头共享同一组KV头。这就是MQA多查询注意力和GQA分组查询注意力。MQA让所有Q头共用一个K头和一个V头KV Cache直接缩小到原来的1/num_heads。比如8个Q头KV Cache就变成原来的1/8约24 MB。代价是单KV头信息量可能不够会影响生成质量。GQA是折中方案把KV头分成若干组每组对应多个Q头。比如8个Q头分成4组每组KV头一个KV Cache缩小到一半质量损失比MQA小很多。MiniMind里最开始我用的是标准MHA训练没问题但推理时在batch8、seq1024会超显存后来改成GQAnum_groups4显存直接降下来40%。改动很小在做attention时先把K、V通过repeat_interleave复制到与Q头数相同再走标准逻辑即可。但要注意repeat_interleave和view可能会带来额外kernel开销最好写成维度广播的形式。3. 省显存训练时如何把显存抠出来3.1 显存都去哪了训练和推理的显存画像完全不同。推理主要是权重 KV Cache 中间激活训练则复杂得多模型权重、梯度、优化器状态、前向激活值还有临时工作区。以AdamW为例每个参数在混合精度训练下大概要占12字节以上权重FP16副本2字节权重FP32主副本4字节动量4字节方差4字节梯度2字节。一个200M参数模型光优化器状态就要800MB。很多人以为模型小就不需要优化显存真训练起来才发现激活值才是大头。每个batch的前向后每层中间输出都要保存下来用于反向传播。以batch_size16、seq_len512、d_model512、num_layers6的MiniMind配置为例单个中间激活是16×512×512×2字节8MB但Transformer里每层有Q/K/V、attention输出、FFN中间两个线性层、两个dropout输出等轻松翻几十倍6层下来激活值能到几百MB甚至1GB以上。3.2 梯度检查点用计算换内存激活值最占用显存但激活值只在前向计算中生成而且在反向传播中会被重新用到。那么能不能不保存全部激活值等反向传播需要时再重新算一遍这就是梯度检查点activation checkpointing也叫重计算。做法很简单把模型按层切分成多个checkpoint段前向时只保存每个段输入到第一层的张量。反向传播时先从checkpoint恢复该段输入然后顺着该段重新执行一遍前向得到后续层的激活值再进行正常反向。这样只需要保存边界处的少量激活内存占用可以降到原来的1/3甚至更低。代价是计算量增加因为每个段的反向都要重跑一次前向。PyTorch中开启方式很简单在Transformer层外面套一个torch.utils.checkpoint.checkpoint(fn, *args)或者用模型对象上的可选参数。MiniMind里我只对最深的6层开启batch从4提到了12训练实际速度慢了约25%但总算能把训练塞进一张8GB显卡。关键心得梯度检查点不是无脑全开就好。浅层模型通常显存压力不大只对后半部分层开启就能获得大部分收益。另外在开启checkpoint时注意不能让输入Tensor被梯度计算截断否则反向时重计算的参数会丢失梯度路径导致训练不收敛。我一开始就栽在这上面梯度全部为NaN后来查了半天才发现是checkpoint函数把输入auto_grad带丢了。3.3 混合精度与低秩微调混合精度训练是另一个“免费午餐”。用FP16或BF16存储激活和做矩阵乘法能减少一半显存和大幅提升速度但要维持主权重以FP32保存避免数值更新太小被低精度吞掉。实际操作时打开AMP的autocast和GradScaler把优化器参数用FP32就行了。BF16的好处是动态范围与FP32一致不会轻易溢出但消费级显卡不一定支持BF16矩阵加速。FP16则需要GradScaler自动缩放loss否则梯度容易下溢。我在MiniMind上用FP16训练时踩过一个坑LayerNorm和embedding在FP16下精度损失过大loss曲线出现异常的尖峰。后来参考常见做法是保留前几层或归一化层在FP32计算其余自动cast问题立刻消失。如果训练的是全新大模型LoRA低秩微调是省显存的终极方案冻结原始权重只训练两个低秩矩阵W W0 BA其中B和A的秩远小于d_model。冻结权重不计算梯度也不在优化器状态中所以显存主要只受batch和激活影响。实测MiniMind在微调文本生成任务时用LoRA让训练显存从6GB降到2GB以内效果还基本持平。这里的核心参数是秩r我一般先用8如果表达能力不够再涨到32rank过大反而过拟合。4. 省时间从注意力算法到工程优化4.1 FlashAttention省显存又省时间传统Attention要构造一个完整的大矩阵S QK^T / sqrt(d)shape是(seq_len, seq_len)所有分数一次性算出来。在长序列下这个矩阵占用的显存和计算量都是O(n^2)的成为整个Transformer中最重的模块。FlashAttention的思路是“既然软最大值要做全局归一化那能不能把计算分块边算边更新最终得到相同结果”它把Q、K、V切分成小块每次只算一个小块的注意力分数并维护当前块的最大值、指数和和归一化分子。当遇到更大分数时可以对之前的结果做重缩放最后不需要存储完整注意力矩阵只保存必要的统计量反向传播时再重算一遍更精简的中间值。效果是双重的显存从O(n^2)降为O(n)速度也因为减少HBM读写而变快。在MiniMind中我没有手写FlashAttention kernel而是直接调用现成的融合算子。但我也做过一个简单的CPU版本作为验证深刻体会到了“分块为什么能省内存”一块一块地进SRAM而不是把整个大矩阵写回显存。4.2 算子融合与自定义kernel很多人在优化LLM时只盯着算法忽略了“kernel launch”和“内存读写”带来的开销。PyTorch里每执行一个算子都会发起一次内核调用并把数据从全局显存读到寄存器或共享内存再写回。Attention模块里如果把softmax、mask、dropout、矩阵乘法拆成一堆独立操作每一层会多几十次内核启动和中间张量读写虽然计算量不变但耗时却可能翻倍。算子融合就是把多个小操作合并成一个内核尽可能在一次读数据中完成所有数学运算。常见例子是“layernorm QKV投影融合”或者“mask softmax dropout 注意力分数融合”。FlashAttention本身就是最典型的算子融合它把QKT、scale、mask、softmax、加权求和全放在一个分块循环里。MiniMind虽然只有6层但把attention内部的几个separate操作改写成融合版本后生成速度提升了约1.8倍。4.3 推理时的内存带宽瓶颈我在优化MiniMind时发现一个反直觉的现象显存占用已经很紧张但实际推理时的瓶颈往往不是计算而是内存带宽。对于单batch生成每个token都要读取模型所有权重200M甚至几个G而计算量相对较小这种情况下GPU的算力是闲置的效率极低。这时候提升吞吐的最好办法是加大batch_size让同一个权重被多个请求共享读取摊薄带宽成本。另一个实用方法是权重量化把FP16权重转成INT8甚至INT4减少权重读取量。比如200M参数从FP16转INT8权重从400MB降到200MB带宽压力减半。MiniMind里我用INT8量化后推理速度提升了约60%生成质量几乎没变化但在小知识密集任务上还是能观察到细微的准确率下降。所以一般对通用场景可以用INT8对需要精确事实的任务保持FP16或混合精度。还有一个容易忽略的点是padding。在批量推理时不同请求的输入长度不一样如果全按批次中最长的序列pad到统一长度会白白浪费大量计算和显存。正确的做法是使用动态padding把相同长度的请求分到一个batch里或者用连续批处理continuous batching让一个batch中某个序列生成完成后立刻插入新任务不让GPU空转。MiniMind的demo虽然用不到这么复杂但我还是实现了长度分桶训练吞吐提升了约30%。5. 深加工隐藏在FFN与归一化里的细节5.1 FFN真正“干活”的模块注意力层是“信息交换所”它让每个token看到上下文中的其他token并聚合信息。但交换完信息之后需要有一个深层处理环节把这些信息转成更有用的表示这个环节就是逐位置的前馈网络FFN。标准Transformer的FFN包含两个线性层中间加一个非线性激活FFN(x) Act(xW1 b1)W2 b2第一个线性层把d_model维映射到4×d_model的中间维度第二个线性层再换回d_model。LLM的参数量中FFN占据了三分之二左右因为中间维度的权重特别大。它本质上是一个“记忆容量池”论文和实验都表明知识大多存储在前馈网络的参数里注意力更多是负责访问与组合。在MiniMind中我一开始为了简单把FFN中间维度从2048砍到了1024结果各种下游任务的loss都略高。后来重新调回2048训练速度慢了40%但Bleu和困惑度都明显改善。这说明FFN不能一味压缩它承载了模型的主要学习容量。5.2 归一化与残差连接的坑没有归一化深层Transformer训练非常容易发散。LayerNorm做的事情很简单对每个样本的特征维度求均值和方差然后标准化再加上可学习的缩放beta和偏移gammay (x - mean) / sqrt(var eps) * gamma beta它保证每个token的输入分布相对稳定不会因为深度增加而剧烈震荡。但很多新手在实现时忽略eps的作用当特征方差很小甚至为0时除以接近0的数会产生无穷大。eps一般取1e-5或1e-6但如果你用fp16训练eps建议取1e-5以上否则数值非常不稳定。MiniMind在后面增加层数时曾因为使用默认1e-6在fp16下产生NaN调大eps才解决。残差连接则让每一层的输出通过加法直接叠加到输入上。这个设计的精妙之处在于即使某一层学不到任何东西模型也能通过残差分支退化成更浅的网络梯度也能绕过权重矩阵直达上游缓解梯度消失。这里有一个容易犯的错归一化放在残差之前还是之后叫pre-norm和post-norm。现在主流LLM都使用pre-norm即先做LayerNorm再做Attention然后残差相加。原因是post-norm在深网络中训练不稳定需要复杂的warmup和超参调整pre-norm的梯度路径更干净训练更稳但理论上表达能力弱一点。5.3 激活函数与扩展方式FFN中用哪个激活函数也很关键。经典的ReLU在Transformer中也能用但现代LLM普遍使用GELU或SiLU也叫Swish。GELU是“概率为φ(x)时保留输入否则置为0”数学形式与ReLU非常接近但更平滑GELU(x) ≈ 0.5x * (1 tanh(sqrt(2/pi) * (x 0.044715x^3)))SiLU更简单SiLU(x) x * sigmoid(x)。它天生就是一个平滑的“非线性门”。更进一步的升级是门控激活函数比如SwiGLU和GeGLU。它们让FFN计算两个分支其中一个分支经过激活后作为“门”来调制另一个分支FFN_SwiGLU(x) (SiLU(xW1) ⊙ (xW3)) W2其中⊙是逐元素乘。由于多了一个参数矩阵实际中间维度通常不是4×d_model而是约2.67×d_model使参数总量与原版FFN持平。MiniMind中我实验过SwiGLU在小模型下收敛速度确实更快但要注意实现时三个矩阵的维度分配xW1和xW3同维度xW2输出回d_model千万别把形状搞错。另外现在很多模型在FFN之后还会加残差连接形成“Attention子层 FFN子层”交替堆叠。我自己的学习体会是深加工模块的设计没有太多花活关键是把归一化位置的逻辑捋清楚保持数值稳定同时按需选择激活函数和中间维度。如果显存不足优先砍FFN的中间维度而不是砍注意力层数这样对模型能力的影响更可控。最后再分享一个小经验很多人把注意力机制背得滚瓜烂熟却在训练时四处碰壁根源就在于位置编码、KV Cache、归一化这些“注意力之外”的细节没有补全。MiniMind这一系列笔记写到这里我自己最大的收获是一个能跑的LLM不是“Attention堆出来的”而是“位置编码 记忆管理 显存优化 时间优化 深加工”这个完整协作系统的产物。如果你也在手写小型GPT建议按上面这五个模块逐个检查把每一步的显存和耗时都拉通测算一遍踩坑才会少一半。
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑