公司动态

环形缓存:滑动窗口注意力在推理引擎中的显存优化利器

📅 2026/9/1 7:02:53
环形缓存:滑动窗口注意力在推理引擎中的显存优化利器
这次的问题可以拆成三句话回答滑动窗口注意力把每个 token 的注意力范围限制在最近 W 个 token 内所以窗口外的 K/V 没有保留价值decode 阶段如果每次滑窗都用数组左移来丢弃最旧 token一次要复制 W×D 个元素成本反而比注意力计算本身还高环形缓存用固定数组加取模定位让新 KV 直接覆盖最旧 KV把写入成本从 O(W×D) 降到 O(D)显存占用也恒定在一个窗口大小。理解这个机制等于同时理解了 KV Cache、滑动窗口注意力以及 vLLM 这类推理引擎里的显存管理思路。这篇文章按“问题背景 → 朴素方案为什么慢 → 环形缓存原理 → 代码实现 → 真实推理栈落地 → 排查清单”这条线展开。如果你是做大模型推理优化、长文本生成、或者准备接 vLLM/TGI 这类服务这个问题的答案值得完整看一遍。先说一个关键结论环形缓存不是“锦上添花”而是滑动窗口注意力在 decode 阶段能落地的必要条件。没有它省下来的显存会被复制开销和内存分配开销重新吃掉。1. 先搞清楚 decode 阶段在解决什么问题1.1 prefill 与 decode 是两个完全不同的阶段大模型自回归生成一次完整推理分成两个阶段。perfill 阶段预填充用户输入一整段 prompt模型一次性并行处理所有输入 token计算出每个 token 的隐藏状态并生成第一个输出 token。这个阶段计算密集GPU 利用率高因为所有输入 token 可以同时参与矩阵乘法。decode 阶段解码模型逐个生成 token每生成一个新 token都要把它和之前所有历史 token 做一次注意力计算然后预测下一个 token。这个阶段是访存密集型的因为每次只算一个 token但需要读取之前全部 token 的 K/V 向量。两个阶段的耗时特征完全不同。prefill 阶段吞吐高decode 阶段延迟敏感。用户能感知到的“一个字一个字蹦出来”主要就是 decode 阶段的时间。1.2 为什么必须有 KV Cache在没有 KV Cache 的朴素实现里每生成一个新 token都要用当前最新的 query重新计算历史所有 token 的 Key 和 Value。这一步的问题是历史 token 的 Key 和 Value 在上一轮解码时已经算过一次重新计算是纯粹的重复劳动。于是标准做法是把历史 token 的 K 向量和 V 向量缓存下来这就是 KV Cache。每次 decode 时新 token 只需要计算自己的 K 和 V。用当前 query 与缓存中历史 token 的 K 做点积。对 scores 做 softmax。用 softmax 权重对缓存中的 V 加权求和。KV Cache 的存在让注意力计算从“所有 token 从头算一遍”变成“只算新 token 的 K/V历史 K/V 直接读取”。这个设计极大降低了 decode 阶段的计算量。1.3 KV Cache 的显存增长是核心矛盾KV Cache 不是静态的它随生成长度线性增长。给一个通用计算公式单个 token 的 KV Cache 大小 2 × num_layers × num_kv_heads × head_dim × 精度字节数其中“2”代表 K 和 V 两份向量。以 Mistral 7B 为例32 层、8 个 KV 头、每个 head 128 维、FP162 字节精度。单 token KV Cache 2 × 32 × 8 × 128 × 2 131072 字节 128 KB这意味着序列长度每增加 1 个 token显存多占 128KB。4096 个 token 时约 512MB100k token 时约 12.8GB。如果并发 100 个请求这个数字会非常夸张。KV Cache 的增长速度直接决定了一个推理服务能支撑多长的上下文、多少并发、多大的 batch。滑动窗口注意力要解决的正是这个增长问题。2. 滑动窗口注意力想解决的问题2.1 核心思想每个 token 只看最近 W 个 token标准自注意力里第 i 个 token 会关注 0 到 i 的所有历史 token计算复杂度是 O(i)。当序列很长时这个复杂度随序列长度线性增长。滑动窗口注意力Sliding Window AttentionSWA把这个范围限制成第 i 个 token 只关注 [i - W 1, i] 这个窗口内的 token窗口大小 W 固定。这样单次注意力计算量从 O(序列长度) 变成 O(W)且 W 是常数。序列再长计算量也不增长。这个概念在 Longformer、BigBird 里被系统化Mistral 7B 也用到滑动窗口注意力窗口默认是 4096。2.2 滑动窗口下旧 token 的 K/V 失去价值这是理解环形缓存的第一个关键点滑动窗口注意力下一个 token 的 K/V 向量有没有价值取决于它还在不在窗口内。当窗口从 [w1, w2, ..., wW] 滑动到 [w2, w3, ..., wW, w(W1)] 时w1 这个 token 已经不再被任何后续 token 关注。它的 K/V 向量之后永远不会再被读取继续保存它只是浪费显存。所以正确的做法是新 token 进入窗口时把最旧 token 的 KV 丢弃。但“丢弃”这件事在工程实现里有很大的性能差异。2.3 滑动窗口能省多少显存配合 KV Cache滑动窗口注意力把缓存上限锁定在固定大小 W。显存占用从“随序列长度线性增长”变成“最多 W 个 token 的 KV”。沿用 Mistral 7B 的 128KB/token 计算序列 4096 token满窗口缓存512MB。序列 100k token全量 KV 缓存是 12.8GB。序列 100k token配合滑动窗口仍然只用 512MB。差距是 25 倍并且随序列继续拉长差距会越来越大。这就是滑动窗口注意力对长文本推理的价值。但这里有一个容易忽略的问题省下来的显存需要用一种廉价的方式把旧 KV 丢掉。如果用最直觉的“数组左移”可能省了显存却把时间花在数据搬迁上。3. 朴素滑窗实现为什么慢3.1 最直觉的做法每次左移一位假设已经有一个长度为 W 的数组 cache里面按顺序存着窗口内 token 的 K/V。当新 token 到来时最直觉的滑动方式是把 cache[1:] 全部复制到 cache[:-1]。把新 token 的 K/V 写到 cache[-1]。写成伪代码def naive_sliding_window(cache, new_kv): # cache: [W, D] 的 KV 数组D 是所有层的 KV 拼接维度 # new_kv: [D] 的新 token KV cache[:-1, :] cache[1:, :] # 整体左移 cache[-1, :] new_kv # 新 KV 写入末尾这个代码逻辑没有错误但性能非常差。问题在于这一行cache[:-1, :] cache[1:, :]这一步是逐元素复制一共要搬动 (W - 1) × D 个数据。3.2 O(W×D) 的隐藏成本解码阶段每生成一个 token都要做一次左移复制。假设模型每层 KV 维度是 D层数是 L那么实际需要移动的元素总量是(W - 1) × D × L把数字代进去看窗口 W 4096。单层单 head 的 KV 维度通常 128层数 32KV 头 8所以 D × L 128 × 8 × 32 32768。每次 decode 需要复制的元素数约为 4095 × 32768 ≈ 1.34 亿个。这个复制量已经接近一次完整注意力计算的 IO 量。更糟糕的是它发生在 decode 的关键路径上直接增加每个 token 的生成延迟。结论是滑动窗口省下的显存如果通过“数组左移”来释放会把时间成本全部吃回去。这显然是不可接受的。3.3 另一个问题反复分配内存如果用更糟糕的方式实现每次滑动都新建一个数组、释放旧数组还会引入额外的内存分配开销。cudaMalloc 和释放显存远慢于常规内存操作在高并发推理里还会造成显存碎片。即使不新建数组如果每次滑动都把整个窗口数据移动到一段新地址同样会有大量 IO。真正的解法是不移动数据只移动“逻辑索引”。4. 环形缓存固定数组 取模定位4.1 数据结构一个固定大小的数组环形缓存Ring Buffer的思路很简单预先分配一块固定长度为 W 的显存/内存不再反复申请和释放。新旧 KV 全部写入这块区域。关键是“写入位置”不固定而是通过绝对位置对窗口大小取模得到物理下标 绝对位置 % 窗口大小举个例子窗口 W 4解码 token 的绝对位置从 0 开始绝对位置 0 → 物理位置 0绝对位置 1 → 物理位置 1绝对位置 2 → 物理位置 2绝对位置 3 → 物理位置 3绝对位置 4 → 物理位置 0覆盖旧位置 0 的 token 0绝对位置 5 → 物理位置 1覆盖旧位置 1 的 token 1窗口内的 token 始终是最近 W 个而物理数组中每个位置都会被反复覆盖。这就是“环形”的含义。4.2 写入新 token一步覆盖写入新 KV 时只需要计算一次取模然后直接覆盖目标位置。伪代码def write_kv(kv_ring, window_size, abs_pos, new_kv): idx abs_pos % window_size kv_ring[idx] new_kv这个操作的时间复杂度是 O(D)也就是只写入新 token 本身的数据不再需要移动任何历史数据。当窗口大小为 2 的幂时可以用位运算优化idx abs_pos (window_size - 1)比如 W 4096那么abs_pos 4095和abs_pos % 4096结果一致但位运算在低层更快。4.3 读取窗口 KV位置映射写入解决了读取也要配套调整。滑动窗口注意力在 decode 时需要读取绝对位置从cur_pos - W 1到cur_pos的所有 KV但这些 KV 物理上可能不连续。读取时需要把窗口内相对位置逐一映射到物理位置def read_window_kv(kv_ring, window_size, cur_pos): start max(0, cur_pos - window_size 1) end cur_pos for abs_pos in range(start, end 1): phys_idx abs_pos % window_size kv kv_ring[phys_idx] # 将 kv 放到与 abs_pos 对应的逻辑位置供注意力计算使用窗口内的排序关系仍然通过绝对位置 abs_pos 维护物理存储顺序可以打乱不影响结果。4.4 环形缓存解决了三个问题第一个问题是数据复制成本。旧 KV 不需要移动新 KV 直接覆盖每次 decode 的缓存写入成本是 O(D)。第二个问题是内存分配。环形缓存只分配一次之后全程复用同一块区域没有反复的 cudaMalloc 和释放。第三个问题是显存上限。整个 KV 缓存的大小被锁定为 W 个 token不管生成多长的序列显存占用都不再继续增长。这三个问题恰好对应工程上的三点要求延迟稳定、无内存碎片、显存可预测。5. 代码对比朴素实现 vs 环形缓存5.1 朴素滑窗版本import numpy as np def naive_generate_with_window(total_tokens, window_size, kv_dim): # 模拟 decode 阶段每步生成一个 token同时滑动窗口 cache np.zeros((window_size, kv_dim), dtypenp.float32) for step in range(total_tokens): new_kv np.random.randn(kv_dim).astype(np.float32) # 滑动左移覆盖丢弃最旧 token 的 KV cache[:-1, :] cache[1:, :] cache[-1, :] new_kv return cache这个版本每步都复制 (window_size - 1) × kv_dim 个元素。total_tokens 越大复制总量越大。5.2 环形缓存版本import numpy as np def ring_generate_with_window(total_tokens, window_size, kv_dim): # 模拟 decode 阶段使用环形缓存存储最近 window_size 个 token 的 KV cache np.zeros((window_size, kv_dim), dtypenp.float32) for step in range(total_tokens): new_kv np.random.randn(kv_dim).astype(np.float32) # 写入绝对位置对窗口大小取模直接覆盖最旧 token 的 KV idx step % window_size cache[idx, :] new_kv return cache这个版本每步只写 kv_dim 个元素不读取、不移动历史数据。5.3 注意力读取的示意代码上面只是 KV 的写入。实际注意力计算时还要按窗口内绝对位置读出 KV 并按顺序做 attention。示意如下import numpy as np def sliding_window_attention_step(query, kv_ring, window_size, cur_pos, head_dim): # query: [head_dim] start max(0, cur_pos - window_size 1) length cur_pos - start 1 scores np.zeros(length, dtypenp.float32) for i, abs_pos in enumerate(range(start, cur_pos 1)): key kv_ring[abs_pos % window_size] scores[i] np.dot(query, key) / np.sqrt(head_dim) scores np.exp(scores - scores.max()) scores scores / scores.sum() out np.zeros(head_dim, dtypenp.float32) for i, abs_pos in enumerate(range(start, cur_pos 1)): value kv_ring[abs_pos % window_size] out scores[i] * value return out这里最核心的一行就是kv_ring[abs_pos % window_size]。物理位置由绝对位置取模得到读取顺序由绝对位置决定。只要这个映射一致无论环形缓存里的数据物理位置多乱最终注意力结果都和标准滑动窗口一致。5.4 定性对比对比项朴素数组左移环形缓存写入耗时O(W × D)O(D)是否移动历史数据是每步移动 W-1 个旧 KV否直接覆盖内存分配次数取决于实现可能反复分配一次性分配显存上限上限固定 W但复制开销大固定 W无额外开销随机读取窗口内 KV物理连续读取直观需要取模映射读取多一次计算如果你的实现允许cache[abs_pos % window_size]这样的取模读取注意在 GPU kernel 里取模运算会带来一点额外指令开销。工程上通常用等价的位运算或固定窗口长度优化掉这部分。6. 真实推理栈里环形缓存怎么落地6.1 HuggingFace transformers 的路线一个常见的认知偏差是HuggingFace transformers 的generate本身不一定使用环形缓存。为了让 Masked Attention 和普通 Attention 统一很多基础实现里 KV Cache 仍然是全量存储的滑动窗口通过 attention mask 实现把窗口外位置的 attention score 设置为 -inf让 softmax 之后权重为 0。这个做法的优势是通用不需要额外处理环形结构缺点是缓存里的旧 token 即使理论不需要实际上仍占着显存。model 参数里即使配置了 sliding_window是否真正释放对应显存还要看具体代码版本和推理 backend。如果你在做一个基于 transformers 库的定制推理脚本想真正省掉窗口外 KV就需要自己实现或切换到 vLLM 这类推理引擎。6.2 vLLM 的 PagedAttention 与 block 级环形管理vLLM 的 PagedAttention 是更工程化的实现。它的核心思路是把 KV Cache 划分成固定大小的 block每个 block 存储若干 token 的 KV通过 block table 管理物理块和逻辑块的映射。在滑动窗口注意力场景下vLLM 会识别到窗口外的 KV 块可以被覆盖复用。这本质上是“block 粒度的环形缓存”不是逐 token 覆盖而是每次把最旧的一个或多个 block 释放回空闲池新 token 从空闲池里申请块。这样做的好处是显存分配单位更大管理开销小。支持多请求并发物理块可以动态分配。滑动窗口与 PagedAttention 的 block table 天然配合。所以如果你在生产环境里用 vLLM 跑 Mistral 这类滑动窗口模型实际享受到的优化正是环形缓存的升级版。6.3 Mistral 与 Longformer 的窗口注意力Mistral 7B 的滑动窗口长度为 4096模型结构上是标准的滑窗注意力。推理时核心问题就回到本文主题整个 KV Cache 如何在不增长的情况下维护最近 4096 个 token 的 KV。Longformer 则更复杂一点它除了滑窗还设计了全局 token 和随机注意力。全局 token 的 KV 不能丢滑窗 token 的 KV 需要按窗口丢弃。在自定义实现里通常需要把全局 token 的 KV 单独存放在另一个缓存区滑窗 token 走环形缓存。这些案例说明环形缓存不一定以“逐 token 环形数组”的形式出现block 管理、分离缓存区本质上都是同一个思想窗口外的 KV 需要被廉价地覆盖掉。7. 影响范围分析它决定了长文本推理的上限7.1 显存从“随序列增长”变成“固定预算”没有环形缓存时滑动窗口只是一个理论省内存方案。一旦实现不到位每次滑窗都要对全窗口做数组平移不仅没省时间还可能让每次 decode 变慢。有环形缓存后模型的 KV Cache 占用被固定在 W × D 上。序列长度 10k 和序列长度 1000w在滑动窗口限制下KV Cache 的内存占用是一样的。这对服务端部署是决定性的你可以提前为每个并发请求预留确定的显存计算最大并发数时只需知道窗口大小不需要预估最坏序列长度。7.2 decode 延迟不再随历史长度增长标准注意力在 decode 时每个 token 要和所有历史 token 计算注意力分数读取 KV 的 IO 量随序列长度线性增加这是长文本生成变慢的一个主要原因。滑动窗口把参与注意力计算的 KV 数量限制在 W 以内。配合环形缓存新 KV 的写入和超窗 KV 的覆盖都是 O(D) 操作不随序列长度增长。于是 decode 速度不会因为“生成了更多 token”而逐步劣化这是长文本生成保持稳定延迟的关键。流式输出场景尤其受益。用户感觉到的“首 token 速度”和“后续 token 速度”可以保持一致不会因为上下文越聊越长而越来越卡。7.3 多轮对话里有一个隐藏权衡这里要说明一个容易踩坑的工程点多轮对话会复用 KV Cache避免重新 prefill 历史对话。但在滑动窗口下如果历史某轮对话的 token 已经滑出窗口它的 KV 已经被覆盖下一次请求就少了一段可复用的缓存。更稳妥的做法是保留多轮对话上下文时要么确保历史 token 仍在窗口内要么接受部分 KV 需要重新 prefill。很多推理框架在面对超长对话时会做“截断 重新 prefill”本质上是把这个问题交给上层调度处理。所以环形缓存不是万能的。它的收益上限是“窗口内的 KV”窗口外的内容本来就不该留在缓存里。7.4 什么时候不需要环形缓存如果模型本身没有使用滑动窗口注意力KV Cache 会一直保留全部历史 token这时候环形缓存没有意义反而会因为覆盖逻辑破坏缓存。如果窗口非常大、接近模型最大上下文长度环形缓存的省内存效果也会减弱。因为本来就是全量差不多够用。环形缓存真正发挥价值的地方在无限长文本生成、流式输出、长对话服务等场景。8. 常见问题与排查方法8.1 位置映射错乱表现生成结果突然错乱尤其是生成较长文本后。原因写入时用了绝对位置取模读取时用了错误的起始位置或者缓存中的相对顺序没有按绝对位置排列。排查打印当前窗口内每个 token 的绝对位置、物理下标、以及实际读出的 KV 来源位置对照检查abs_pos % window_size是否一致。8.2 attention mask 没有跟着窗口滑动表现注意力结果和全量注意力偏差越来越大或者模型重复生成。原因cache 已经实现环形覆盖但 attention mask 仍然是全量掩码导致窗口外的 token 仍然参与注意力计算。此时缓存里的 KV 可能已经被覆盖成新 token却仍然用旧 mask 计算数据和掩码错位。排查把 mask 打印出来检查 mask 覆盖的 token 范围是否随cur_pos - W 1动态变化。8.3 多轮对话缓存复用失败表现第二轮对话内容重复、混乱或者前面的上下文被“遗忘”得比预期更严重。原因第一轮生成的 KV 在第二轮开始时可能已经滑出窗口并被覆盖但上层逻辑仍然认为 KV 缓存可用没有重新 prefill 被覆盖的历史。排查在多轮切换时查看当前窗口覆盖的最小绝对位置如果它大于第二轮需要的起始上下文就需要重新 prefill。8.4 窗口大小不是 2 的幂时性能下降表现取模运算耗时偏高或者 GPU kernel 里出现额外的分支。排查将窗口大小设置为 2 的幂例如 4096、8192用位运算替代取模。如果必须使用非 2 的幂窗口要评估取模开销是否可接受。8.5 与 FlashAttention 的兼容问题表现attention 计算很快但 KV 缓存写入和读取逻辑不一致最终结果错误。排查FlashAttention 有自己的分块和 IO 策略滑动窗口的 mask 仍要正确传给 kernel。实现时先在 CPU 上用小规模数据跑通无 FlashAttention 版本再切换到 Flash 版本做一致性对比。8.6 环形缓存里读取到的历史 KV 是错的表现缓存大小正确但输出质量下降或者 attention score 数值异常。原因写入覆盖没有严格按照绝对位置取模或者 prefill 阶段的 KV 起始位置和 decode 阶段的绝对位置没有对齐。排查在 prefill 阶段记录第一个 token 的绝对位置decode 阶段从同一个基准继续累加绝对位置而不是从 0 重启。问题现象可能原因排查方式解决方案生成结果变错乱位置映射不一致打印绝对位置与物理下标统一abs_pos % window_size映射注意力范围不对attention mask 未跟随窗口检查 mask 覆盖范围动态更新 mask 为最近 W 个 token多轮对话上下文丢失窗口外 KV 被覆盖后仍按可复用处理检查当前窗口起始位置覆盖部分重新 prefill取模成为性能瓶颈窗口大小非 2 的幂profile kernel 耗时窗口设为 2 的幂并用位运算FlashAttention 结果不一致滑动 mask 未传入或传错CPU 与非 Flash 版本对比检查 kernel 的 window 参数9. 小结与后续可扩展方向滑动窗口注意力在 decode 阶段使用环形缓存本质上是解决一个工程矛盾窗口外的 KV 必须丢弃但丢弃方式不能引入与窗口大小成正比的复制开销。环形缓存用固定数组加取模定位让写入、读取、覆盖都变成常数时间操作同时把显存占用限制在一个窗口大小内。最先应该验证的内容有几个第一在 CPU 小规模环境下跑通环形缓存的读写映射确认和标准滑动窗口结果一致第二在 GPU 上用性能分析工具对比数组左移和环形缓存两种实现的 decode 耗时第三检查多轮对话场景下 KV 复用失败时是否出现位置错乱。最容易踩坑的地方集中在三处绝对位置和物理下标的映射不一致、attention mask 未跟随窗口滑动、多轮对话缓存复用时的重新 prefill 逻辑。这三类问题都属于“逻辑正确但实现错位”调试时一定要把每个 token 的绝对位置和物理下标打出来对比。后续可以继续扩展的方向包括把环形缓存与 PagedAttention 的 block 管理结合在服务端实现高并发长对话研究 StreamingLLM 提出的 attention sink 思路在窗口外保留少量初始 token 以稳定长文本生成以及把滑动窗口和 FlashAttention 的分块策略进一步融合减少 kernel 内外存搬运。分享一个最直接的结论如果你的推理引擎已经支持滑动窗口注意力不要自己去用数组左移实现如果你在做自定义推理 kernel环形缓存应该是默认选择而不是优化项。