最近在优化长文本生成推理时反复琢磨一个有意思的问题为什么实现滑动窗口注意力时decode 阶段要专门引入环形缓存网上资料大多只讲“窗口限制注意力范围”却很少讲清楚“缓存如何组织”。本文从标准自注意力出发一步步推导滑动窗口注意力的 decode 流程重点拆解环形缓存的设计思路、最小实现与工程坑点顺便把 KV Cache、位置编码、常见报错这些关联知识也一并梳理清楚。1. 从标准自注意力说起1.1 自注意力的计算过程Transformer 的核心是自注意力Self-Attention机制。给定输入序列 (x_1, x_2, \dots, x_T)模型会为每个 token 生成对应的查询向量 (Q)、键向量 (K) 和值向量 (V)然后通过下面的公式计算输出[ \text{Attention}(Q,K,V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V ]其中 (d_k) 是键向量的维度除以 (\sqrt{d_k}) 是为了缩放点积避免数值过大导致 softmax 梯度消失。这个过程看起来不复杂但有一个容易被忽略的点在自回归生成场景下标准注意力默认每个 token 都能看到序列中所有位置的信息。也就是说如果当前生成了 1000 个 token下一个 token 的注意力计算会完整地扫描这 1000 个 token 的 (K) 和 (V)。直接按这个公式做推理存在两个问题计算量会随着序列长度线性增长显存中需要保存所有历史 token 的 (K)、(V)而且不断增长。所以在实际推理框架中几乎所有实现都会引入一个叫 KV Cache 的优化把已经计算过的 (K)、(V) 缓存下来避免每个新 token 都重新计算整个序列的 (K)、(V)。1.2 KV Cache 让 decode 更高效自回归生成通常分为两个阶段prefill预填充阶段一次性处理完整的输入 prompt并行计算出所有 prompt token 的 (K)、(V)并写入缓存decode解码阶段逐 token 生成每一步只处理最新生成的那个 token把它的 (K)、(V) 追加到缓存末尾然后只查询缓存中的历史 (K)、(V)。KV Cache 带来的收益非常明显如果缓存长度为 (L)标准 attention 计算只需要计算新 token 的 (Q) 与缓存中所有 (K) 的点积不需要重复计算历史 token 的 (K)、(V)。这里贴一段简化版的标准 KV Cache 追加逻辑import torch class SimpleKVCache: def __init__(self, max_len, num_heads, head_dim): self.k_cache torch.zeros(max_len, num_heads, head_dim) self.v_cache torch.zeros(max_len, num_heads, head_dim) self.valid_len 0 def append(self, key, value): self.k_cache[self.valid_len] key self.v_cache[self.valid_len] value self.valid_len 1在标准 KV Cache 里缓存长度等于已经生成过的 token 总数。对于短序列来说这很轻松但对于长文本生成KV Cache 的显存开销会迅速膨胀。1.3 滑动窗口注意力限制视野而不是重算历史为了控制长文本带来的显存和计算压力滑动窗口注意力Sliding Window Attention被提了出来。它的核心思路很简单每个 token 不再看整个序列而是只看它前面最近的 (W) 个 token其中 (W) 是窗口大小。比如 (W4)生成第 100 个 token 时它只 attend 第 96 到第 99 个 token更早的 token 虽然依旧存在但不再参与当前步的注意力计算。这种设计在很多大模型里都能看到。Mistral 系列模型就采用了滑动窗口注意力Gemma 2 也用了局部注意力和全局注意力交替的变体。它解决的问题很直接把注意力计算复杂度从 (O(L^2)) 降到 (O(L \cdot W))其中 (L) 是序列长度(W) 是窗口大小理论上离得越远的 token 对当前生成的影响越小限制窗口不会有太大效果损失让模型推理时的显存开销变得可控不再随序列长度无限增长。但这里有一个关键坑点滑动窗口限制的是“注意力参与范围”并不代表 KV Cache 会自动减少。如果你只是把 attention mask 改成窗口形式KV Cache 仍然会不断追加新 token缓存长度依旧线性增长。这也是环形缓存要解决的核心问题。2. decode 为什么需要环形缓存2.1 prefill 与 decode两种完全不同的计算模式要理解为什么 decode 需要环形缓存先要知道 prefill 和 decode 的区别。在 prefill 阶段prompt 的所有 token 是同时已知的可以并行计算。此时即使使用滑动窗口每个位置也能一次性算出注意力输出KV Cache 可以一次性构建完整。但在 decode 阶段模型只能一个一个地生成 token每个新 token 的 (K)、(V) 都会追加到缓存中。如果此时缓存依然采用“无限追加”的模式那么即使窗口只关注最近 (W) 个 token缓存的长度仍然会持续增长。我们假设窗口大小是 (W2048)但生成长度达到 20000 token。如果按朴素的窗口 attention 实现每一步计算时只有最近 2048 个 token 参与 softmax但 KV Cache 中仍然保存了 20000 个 token 的 (K)、(V)显存占用和计算量并没有真正降下来。换句话说窗口限制只影响了“读哪些数据”没有影响“存哪些数据”。这在长文本场景下依然会爆炸尤其是大模型的层数多、头数多、维度大时。2.2 朴素滑动窗口实现的内存问题我们算一笔简单的账。假设模型有 32 层32 个注意力头每个 head 维度是 128使用 FP16 存储。每个 token 的 KV Cache 大小大约是[ 32 \times 32 \times 128 \times 2 \times 2 \text{ bytes} \approx 524288 \text{ bytes} ]也就是 0.5 MB 左右。如果序列长度是 10000KV Cache 就大约需要 5 GB。再叠加多 batch、多个并发请求显存很快会被占满。而滑动窗口希望通过“只保存最近 (W) 个 token 的 KV”把缓存大小固定在一个常数范围内。这样无论生成多长显存占用都不会随生成长度线性增长。于是一个很自然的需求出现了当 KV Cache 满了以后需要把最旧的 token 覆盖掉给新 token 腾出位置。2.3 环形缓存的核心思想环形缓存Ring Buffer / Circular Buffer是一种固定容量的数据结构。它有两个核心指针写入指针write index下一个新数据要写入的位置有效长度valid length当前缓存中有效数据的数量。当写入指针到达缓存末尾时它会回到开头继续写入像一个环一样循环。在这个过程里最早写入的数据会被新数据覆盖。对应到滑动窗口注意力中缓存容量固定为窗口大小 (W)生成一个新 token 时把它的 (K)、(V) 写入当前 write_idx 指向的位置write_idx 向前移动一位如果已经到末尾则回到 0当有效长度达到 (W) 后每次写入都会覆盖掉最旧的一个 token。这样做的好处是缓存占用显存是固定的不随生成长度增长覆盖旧数据的时间复杂度是 (O(1))只需要改变指针注意力计算时始终只读取缓存中最近 (W) 个 token 的 (K)、(V)正好对应滑动窗口的语义。下图用 ASCII 简图表示环形缓存覆盖过程。假设缓存容量为 4已经写入了 token 0、1、2、3写入指针回到位置 0初始状态 slot: [0] [1] [2] [3] pos: 0 1 2 3 write_idx - 0 写入 token 4 覆盖 slot 0token 4 进入 slot: [4] [1] [2] [3] write_idx - 1当继续写入 token 5 时会覆盖 slot 1缓存里始终保留最近 4 个 token。2.4 为什么是“环形”而不是删除旧数据有同学会问既然只需要最近 (W) 个 token那直接把最旧的数据删掉不就行了为什么还要设计成环形如果你用普通数组来管理删除最旧的数据通常意味着要把后面的数据整体往前搬移这是一个 (O(W)) 的操作在 GPU 上会造成不必要的访存开销和性能抖动。环形缓存不一样。它不搬移任何数据只通过指针跳转来实现“逻辑上删除了最旧数据物理上新数据直接覆盖旧位置”。这个特性在 GPU 上非常友好因为我们可以提前分配一整块连续显存然后像流水线一样不断覆盖写入。另外环形缓存也天然适配 batch 场景。多个请求的 KV Cache 可以各自维护独立的 write_idx避免了频繁分配和释放显存导致的内存碎片。3. 环形缓存的关键设计点3.1 缓存容量与窗口大小的关系环形缓存的容量需要和窗口大小匹配。最简单的设计是[ capacity window_size ]这样缓存满的时候刚好保存最近 (W) 个 token。如果容量小于窗口大小一些本应参与注意力的 token 会被过早覆盖导致生成质量下降如果容量大于窗口大小则需要额外记录每个 token 的全局位置在算注意力时把窗口之外的数据过滤掉实现复杂度会上升。实际工程中有些框架会把容量设置成 block 的整数倍比如每 16 个 token 作为一个 block这样便于做内存对齐也能减少索引计算的开销。3.2 索引计算write_idx 与取模环形缓存最核心的索引计算有两个写入位置write_idx (write_idx 1) % capacity某个全局位置对应的槽位slot global_pos % capacity这里要特别注意全局位置global position和缓存槽位slot index不是同一个概念。token 的全局位置是它在整个序列中的绝对位置比如第 100 个 token 的 global_pos 是 99而它的缓存槽位是 (99 \mod capacity)。如果只看 write_idx 而不记录全局位置很容易在学生时代常见的环形队列作业里踩坑因为覆盖之后旧的全局位置信息不会自动消失所以必须额外维护一个数组或指针来记录“哪些槽位是有效的”以及“每个槽位对应的全局位置是多少”。3.3 位置编码不能把槽位当成位置环形缓存最隐蔽的坑点在于位置编码。如果模型使用绝对位置编码比如 learned absolute position embedding那么每个 token 的位置信息已经被加到输入中(K)、(V) 里实际上已经包含了位置信息环形缓存只要原样保存 (K)、(V) 即可。但如果使用相对位置编码比如 GPT-NeoX 和 LLaMA 系列使用的 RoPERotary Position Embedding情况就复杂了。RoPE 的有效性来自 query 和 key 之间的相对位置差[ \text{score}(q_m, k_n) q_m^T R_{m-n} k_n ]这意味着在计算注意力分数时必须知道 query 对应当前 token 的全局位置 (m)以及被 attend 的 key 对应的全局位置 (n)然后根据两者之差来旋转位置编码。这就带来一个硬性要求环形缓存的每个槽位必须额外保存它对应的全局位置。否则你无法知道缓存里的 key 到底属于第几个 token也无法正确计算 RoPE。很多实现里会维护一个position数组或者start_idx来标记窗口最早 token 的位置。每次 attention 时根据全局位置重新生成 RoPE 的 cos/sin 值而不是直接使用缓存里的固定位置。3.4 与标准 KV Cache、PagedAttention 的关系从数据结构的角度看标准 KV Cache 其实可以理解成一个“无限容量”的线性写入缓存而环形缓存是固定容量的覆盖写入缓存。两者并不冲突。在实际推理框架中vLLM 等方案又进一步引入了 PagedAttention把 KV Cache 按固定大小的 block 管理。每个 block 内部连续存储block 之间通过索引连接。滑动窗口 环形缓存的思路天然可以和张量并行、分页管理结合起来。理解这些概念之间的关系后续读推理框架源码时就不会觉得混乱。4. 一个最小可运行的环形缓存示例下面我们实现一个极简的环形缓存 demo用 PyTorch 张量模拟 decode 阶段。4.1 环形缓存类设计import torch class RingKVCache: 简易环形 KV 缓存用于滑动窗口注意力 decode 阶段。 capacity 表示缓存槽位数量window_size 表示滑动窗口大小。 def __init__(self, capacity, num_heads, head_dim, window_sizeNone): self.capacity capacity self.num_heads num_heads self.head_dim head_dim self.window_size window_size if window_size is not None else capacity self.k_cache torch.zeros(capacity, num_heads, head_dim) self.v_cache torch.zeros(capacity, num_heads, head_dim) # slot_global_pos 记录每个槽位对应的全局位置-1 表示无效 self.slot_global_pos torch.zeros(capacity, dtypetorch.long) - 1 self.write_idx 0 self.valid_len 0 self.max_global_pos -1 def append(self, key, value): 写入一个新的 KV。 key/value 形状: [num_heads, head_dim] new_global_pos self.max_global_pos 1 self.k_cache[self.write_idx] key self.v_cache[self.write_idx] value self.slot_global_pos[self.write_idx] new_global_pos self.max_global_pos new_global_pos self.write_idx (self.write_idx 1) % self.capacity if self.valid_len self.capacity: self.valid_len 1 def attention(self, query, scale1.0): 计算当前 query 在滑动窗口内的注意力输出。 query 形状: [num_heads, head_dim] 返回形状: [num_heads, head_dim] # 有效 mask只有已经写入过数据的槽位才有效 valid_mask self.slot_global_pos 0 # 如果当前有效 token 数量超过 window_size只保留最近 window_size 个 if self.valid_len self.window_size: min_global_pos self.max_global_pos - self.window_size 1 valid_mask self.slot_global_pos min_global_pos # 扩展维度方便矩阵乘法 q query.unsqueeze(0).unsqueeze(0) # [1, 1, num_heads, head_dim] k self.k_cache.unsqueeze(0).permute(0, 2, 1, 3) # [1, num_heads, capacity, head_dim] v self.v_cache.unsqueeze(0).permute(0, 2, 1, 3) # [1, num_heads, capacity, head_dim] # scores 形状: [1, 1, num_heads, capacity] scores torch.matmul(q, k.transpose(-2, -1)) * scale # 构造 mask形状需要对齐 scores 的最后两维 mask ~valid_mask mask mask.unsqueeze(0).unsqueeze(0).unsqueeze(0) # [1, 1, 1, capacity] scores scores.masked_fill(mask, float(-inf)) probs torch.softmax(scores, dim-1) out torch.matmul(probs, v) # [1, 1, num_heads, head_dim] return out.squeeze(0).squeeze(0)这段代码有几个设计要点slot_global_pos是核心。它记录了某个槽位里保存的是第几个 token这是位置编码和多轮覆盖时判断有效性的关键。write_idx只负责“下一次写入到哪里”不直接参与 attention 计算。valid_mask一开始用slot_global_pos 0判断槽位是否有效当有效 token 数量超过窗口大小时再根据全局位置过滤掉更早的 token。4.2 模拟 decode 过程下面我们用随机数据模拟一个生成过程观察缓存状态如何变化。capacity 4 num_heads 2 head_dim 8 cache RingKVCache(capacity, num_heads, head_dim, window_size4) for step in range(6): # 模拟当前新 token 的 K/V 和 query key torch.randn(num_heads, head_dim) value torch.randn(num_heads, head_dim) query torch.randn(num_heads, head_dim) cache.append(key, value) out cache.attention(query, scale1.0 / (head_dim ** 0.5)) print(fstep{step}, out_shape{tuple(out.shape)}, fvalid_len{cache.valid_len}, write_idx{cache.write_idx}, fglobal_pos{cache.slot_global_pos.tolist()})预期输出大致如下step0, out_shape(2, 8), valid_len1, write_idx1, global_pos[0, -1, -1, -1] step1, out_shape(2, 8), valid_len2, write_idx2, global_pos[0, 1, -1, -1] step2, out_shape(2, 8), valid_len3, write_idx3, global_pos[0, 1, 2, -1] step3, out_shape(2, 8), valid_len4, write_idx0, global_pos[0, 1, 2, 3] step4, out_shape(2, 8), valid_len4, write_idx1, global_pos[4, 1, 2, 3] step5, out_shape(2, 8), valid_len4, write_idx2, global_pos[4, 5, 2, 3]从输出可以看到前 4 步缓存按顺序填满第 4 步开始写入位置回到 slot 0把旧的 token 0 覆盖成 token 4第 5 步写入 slot 1把 token 1 覆盖成 token 5缓存中始终保留最近 4 个 token 的 KV 数据。4.3 运行结果说明这个 demo 只保留了滑动窗口注意力最核心的缓存逻辑并没有包含多头线性投影位置编码多层叠加批量推理。但在工程理解上已经足够说明问题。真正的推理框架里每一步不只是改写write_idx还需要处理 batch 中每个序列不同的valid_len。这也是很多长文本推理优化的难点所在。5. 工程实现中的常见问题与排查5.1 常见问题速查表问题现象常见原因解决思路生成结果逐渐错乱甚至出现乱码位置编码未按全局位置计算保存slot_global_pos或start_idxRoPE 按相对位置计算缓存里出现大量无效数据softmax 时概率异常mask 构造错误把无效槽位当成有效位置打印slot_global_pos和write_idx对照检查显存没有下降capacity 仍然等于最大序列长度没有真正限制 KV 缓存大小将 capacity 设为窗口大小并确认所有层共用同一套设置长文本越往后效果越差模型依赖早期 token窗口内没有足够上下文考虑保留 attention sinks参考 StreamingLLM 思路decode 性能比预期低每次 attention 都重新生成 mask 和位置编码预分配 mask、预计算 RoPE 的 cos/sin多 batch 时一个序列出错影响其他序列write_idx 和 valid_len 是全局共享的每个序列单独维护一份缓存状态5.2 数据覆盖导致结果错乱很多同学第一次写环形缓存时会遇到“覆盖后旧数据没有清掉”的困惑。比如容量为 4写入 5 个 token 后slot 0 保存的是 token 4slot 1 保存的是 token 1。如果你在 attention 时仍然按照“从 0 到 valid_len-1 都有效”的逻辑取数就会把 token 1 误认为是序列中第二个 token位置信息全部错乱。排查思路打印slot_global_pos确认每个槽位对应的全局位置打印write_idx确认覆盖写入都是从最旧位置开始的在 attention 的 mask 构造处加断言确保 mask 和实际全局位置一致。5.3 位置编码导致长文本效果下降如果是绝对位置编码问题通常不大因为位置信息在进入模型前就已经叠加。但使用 RoPE 等相对位置编码时位置信息是在注意力分数计算阶段动态注入的这意味着每次计算score时都要知道 key 的全局位置。如果你偷懒直接用slot_idx代替global_pos在缓存未覆盖阶段可能是正确的但一旦发生覆盖slot_idx和global_pos就完全对不上了。最终表现就是短文本正常、长文本质量断崖式下降。正确做法是在缓存中额外维护slot_global_pos在 attention 计算前把每个 key 对应的全局位置取出来然后按相对位置生成旋转矩阵。5.4 prefill 与 decode 缓存不统一另一个常见问题是 prefill 阶段使用普通线性缓存decode 阶段直接切换成环形缓存导致两边的 KV 拼接对不上。最直接的解决方式是从一开始就统一缓存接口。prefill 阶段也可以用环形缓存把 prompt 中前 (W) 个 token 依次写入环形缓存然后从第一个新 token 开始继续使用同样的 append 逻辑。这样 prefill 和 decode 的代码路径完全一致后期排查也方便。5.5 多 batch 场景的索引竞争在线推理一般会同时处理多个请求。如果所有序列共享一个write_idx那么一个序列的写入会立刻覆盖另一个序列的数据结果必然错乱。正确做法是为每个序列维护独立的write_idx、valid_len、slot_global_pos。在 batch 推理时可以把这些状态打包成张量用向量化操作同时更新多个序列。6. 工程最佳实践与优化建议6.1 缓存容量固定避免动态扩容环形缓存最大的优势之一就是“固定容量、零搬移”。在工程实现中尽量在初始化阶段就把容量确定下来避免在推理过程中动态扩容。动态扩容涉及重新申请显存把旧数据拷贝到新缓冲区更新所有指针。这些操作在长文本并发推理场景下会带来明显的延迟抖动而且容易造成显存碎片。更好的方式是根据最大并发量和窗口大小在初始化时一次性把显存预留好。6.2 用独立数组保存全局位置强烈建议在缓存结构里单独维护一个slot_global_pos数组。不要试图从write_idx和valid_len反推每个槽位的全局位置因为在覆盖发生之后这种推算很容易出错。有了全局位置数组之后无论是调试、生成可视化还是实现 RoPE都会方便很多。它可以被看作是一份“元数据”随 KV 一起更新。6.3 保留 attention sinks提升长文本效果滑动窗口注意力在理论上很优雅但实际使用时有一些模型会出现“越长越笨”的问题。一个重要原因是某些位置对很多 token 都有全局影响被称为“注意力汇点”attention sinks。典型例子是序列开头的第一个 token或者分隔符 token。StreamingLLM 论文提出一个很实用的思路在滑动窗口之外单独保留最开始的几个 token 的 KV。这样既控制了缓存大小又不会丢失关键的全局信息。如果你在长文本场景下发现窗口注意力效果下降严重可以考虑保留少量 attention sink token。6.4 优先复用成熟推理框架如果你是业务开发不一定要从零实现环形缓存。vLLM、TensorRT-LLM、Hugging Face Transformers 等框架已经实现了比较成熟的 KV Cache 管理方案并且支持滑动窗口注意力。自己实现时容易遇到张量并行时缓存切分问题不同模型的 RoPE 实现差异动态 batch 和抢占调度。所以建议先把框架自带能力用起来只有当框架不满足需求时再基于源码做二次开发。自己写 demo 用于理解原理用框架跑生产环境这可能是更稳妥的组合。6.5 性能观测与调优指标优化环形缓存时至少关注三个指标KV Cache 显存占用确认它是否随生成长度保持常数而不是线性增长单 token decode 延迟观察引入环形缓存后单步生成延迟是否稳定有效 token 命中率统计 attention 中实际参加计算的 token 数与缓存容量之比。跑性能测试时建议用模型的真实配置层数、头数、维度压测而不是只跑小规模 demo。很多时候优化效果要在显存压力较大时才能体现出来。7. 总结与延伸学习把滑动窗口注意力和环形缓存放在一起看核心其实只有两件事限制注意力的计算范围控制 KV Cache 的存储范围。窗口大小决定了模型能看到多远环形缓存则决定 KV Cache 有多大。两者配合才能让长文本生成在显存可控的前提下保持稳定推理。如果要从这个方向继续深入下一步可以关注几个关联主题StreamingLLM理解 attention sinks学习如何让滑动窗口在无限长文本中保持效果PagedAttention / vLLM学习 KV Cache 的分块管理理解为什么 block 管理在并发推理中更有优势KV Cache 量化研究把 KV Cache 从 FP16 压缩到 INT8 甚至更低精度进一步降低显存占用RoPE 的工程实现在手写环形缓存时正确实现相对位置编码往往是最大的难点。如果本文对你有帮助可以收藏备用。也建议你拿一个小的 decoder-only 模型把上面的 RingKVCache 集成进去跑一遍长文本生成这样对“为什么 decode 要用环形缓存”的理解会比只看文章深刻很多。