最近在调一个长文本生成服务时,遇到一个非常典型的显存问题:模型启动后生成速度正常,一旦对话轮次变多、上下文变长,GPU 显存占用几乎线性上升,最后直接 OOM。第一反应是 batch 开太大,把 batch 调小后依然如此;后来逐层打印显存占用,才发现问题出在 KV cache 上——它不仅没有复用旧空间,还在 decode 阶段不断追加新缓存,把滑动窗口注意力的优势完全浪费了。
这篇文章围绕一个核心问题展开:为什么滑动窗口注意力(Sliding Window Attention,SWA)在 decode 阶段要使用环形缓存(Ring Buffer / Circular Buffer)。我会先讲清楚滑动窗口注意力、decode 阶段、环形缓存三者的概念,再拆解环形缓存解决的内存复用、位置索引、并发访问等问题,最后通过一个可运行的 Python 示例演示完整实现,并整理生产环境中的典型踩坑点。
读完这篇文章,你应该能回答下面几个问题:
- 滑动窗口注意力的窗口是怎么“滑动”的;
- 为什么 decode 阶段不能用普通 KV cache 无脑累积;
- 环形缓存为什么能把显存占用锁死在固定大小;
- 位置编码(尤其是 RoPE)为什么是环形缓存实现的关键难点。
1. 背景与核心概念
1.1 从一次长文本生成显存暴涨说起
在自回归大模型的推理过程中,模型每生成一个 token,都要让当前 token 的 Query 与之前所有 token 的 Key、Value 做注意力计算。如果每次生成都重新计算所有历史 token 的 K/V,计算量会随序列长度平方级增长,完全不可接受。所以业界普遍采用 KV cache:在第一次遇到每个 token 时,把它对应的 Key 和 Value 存入显存,后续生成时直接复用,避免重复计算。
KV cache 的问题也随之而来:它占用显存的大小与序列长度成正比。序列越长,缓存越大。如果模型上下文窗口是 32K、64K,而服务的并发又高,显存会迅速被吞掉。滑动窗口注意力可以限制每个 token 的注意力范围,自然也应该限制 KV cache 的大小。但如果实现时仍然用普通追加式缓存,滑动窗口只限制了“计算范围”,没有限制“存储范围”,显存问题仍然存在。
1.2 滑动窗口注意力是什么
标准自注意力中,序列里第i个 token 需要计算与所有0..i-1位置 token 的注意力分数。当序列长度为n时,注意力矩阵大小是n*n,复杂度是 O(n²)。这既拖慢训练,也让推理时每一步需要读取大量历史 KV。
滑动窗口注意力做了简化:每个 token 只关注它前面最近W个 token,其中W是窗口大小。以W=4为例:
- 第 5 个 token 只看第 1 到第 4 个 token;
- 第 6 个 token 只看第 2 到第 5 个 token;
- 第 7 个 token 只看第 3 到第 6 个 token。
窗口每前进一个 token,最前面的旧 token 就会被“滑出”窗口。这种设计让注意力复杂度从 O(n²) 降为 O(n·W)。当W远小于n时,计算量大幅下降。Longformer、Mistral、Qwen 等模型架构在部分层中采用了类似的局部注意力思想。
1.3 decode 阶段到底做了什么
大模型推理通常分成两个阶段:prefill 阶段和 decode 阶段。
- prefill 阶段:把用户输入的 prompt 一次性编码,计算每个输入 token 的 K/V,并缓存下来。这个阶段并行度高。
- decode 阶段:模型逐 token 生成输出。每生成一个新 token,把新 token 当作当前 Query,与缓存中的 K/V 做注意力计算,算出下一个 token 的概率分布,然后采样,再把新 token 追加到序列中,更新 KV cache。
decode 阶段的特点是“一次只前进一个 token”。它天然适合增量计算:历史 token 的 K/V 不需要重算,只需把新 token 的 K/V 追加到缓存中。而“追加”这个动作,正是环形缓存要优化的对象。
1.4 环形缓存是什么
环形缓存也叫循环缓冲区,是一种固定大小的数据结构。它与普通数组/列表最大的区别在于:当缓冲区写满后,新数据会写入“最旧数据”的位置,从而覆盖旧数据。它通过头指针和一个模运算实现循环写入,写入和读取时间复杂度都是 O(1)。
举个例子,容量为 4 的环形缓存写入 A、B、C、D 后已经写满,此时再写入 E,会覆盖 A 的位置;继续写入 F,覆盖 B 的位置。缓冲区中始终保留最近写入的 4 个元素。
环形缓存非常适合描述“滑动窗口”的语义:窗口滑走的数据不需要显式销毁,直接让新数据覆盖即可。
2. 滑动窗口、KV cache 与环形缓存的关系
2.1 传统 KV cache 的线性增长问题
不使用滑动窗口时,所有历史 token 的 K/V 都要保留。假设模型层数为L,每层 KV cache 大小约为:
2(K 和 V)× L × num_heads × head_dim × seq_len × 每个元素字节数当seq_len从 1024 涨到 8192,KV cache 也涨 8 倍。如果是 7B 模型、32 层、4096 隐藏维度、又开了多 batch,显存很容易被缓存占满。因此长文本推理场景中,KV cache 的管理是性能优化的核心。
2.2 滑动窗口注意力如何影响历史 KV
滑动窗口注意力相当于给每个 token 规定了“有效视野”。窗口滑到第t个 token 时,只有位置在t-W+1到t之间的 K/V 会被用到。比t-W+1更早的 K/V,从注意力计算角度已经彻底失去作用。
这带来一个重要推论:在 decode 阶段,只要保证缓存中保存最近 W 个 token 的 K/V,就足以计算结果。再早的 K/V 可以安全丢弃。
问题变成了:如何高效地“丢弃旧缓存、写入新缓存”?如果每步都通过移动数组元素来删除头部、追加尾部,即把后面的元素全部前移,会产生 O(W) 的搬运开销。而环形缓存通过覆盖写实现 O(1) 的淘汰和插入,正好匹配这个需求。
2.3 环形缓存解决的问题
综合来看,环形缓存解决的是三个问题:
- 内存固定:缓存数组预分配 W 个位置,不随序列长度增加而扩容,显存占用可控。
- 淘汰高效:写满后自动覆盖最旧数据,不需要移动元素,也没有额外释放/申请开销。
- 语义匹配:环形缓存的“滚动覆盖”天然对应滑动窗口的“窗口平移”。
因此,滑动窗口注意力在 decode 阶段使用环形缓存不是技巧性优化,而是逻辑上的必然选择:窗口滑到哪,缓存就跟到哪。
3. 为什么 decode 阶段必须使用环形缓存
3.1 窗口滑动的本质是“遗忘”
很多初学者以为滑动窗口只是“注意力掩码不同”,掩码里把窗口外位置置为负无穷即可。这在 prefill 阶段是对的,但在 decode 阶段会带来浪费。
考虑一个窗口大小W=4的场景。当生成到第 10 个 token 时,它只需要第 7、8、9、10 个 token 的 K/V。第 1 到第 6 个 token 的 K/V 已经永远不会被后续任何 token 使用。如果还把它们的缓存保存在显存中,相当于让模型持续占用一块“永远不会被读取”的内存。
正确的做法是:窗口每滑动一步,就允许新 K/V 覆盖最旧 K/V。这正是环形缓存的机制。所以从工程角度看,滑动窗口注意力的 decode 实现,本质上就是一个基于环形的滑动 KV cache。
3.2 普通数组覆盖会丢失逻辑位置
有人会问:我不用环形缓存,直接用普通数组,新 K/V 写到 old slot 的行不行?
可以写,但要注意逻辑位置和物理位置的对应关系。假设用普通数组cache[0..W-1]存 K/V,并规定“新 token 永远写到cache[step % W]”。这个写法和环形缓存其实没有本质区别,只是缺了一个明确的结构化封装。
更关键的问题是:注意力计算时必须知道每个 K/V 对应的原始位置 ID。
如果模型使用绝对位置编码,那么 K/V 所在的物理槽位一旦被覆盖,新的 K/V 要是沿用了旧的物理位置号,位置语义就错了。如果模型使用 RoPE 这类相对位置编码,则需要根据“当前 Query 的原始位置”和“缓存里 K/V 的原始位置”计算旋转角度。因此,无论哪种位置编码,缓存中都要额外保存每个 K/V 的原始 position id,不能只看物理数组下标。
环形缓存的价值就在于,它让“物理位置”和“逻辑位置”解耦:物理位置固定为 W 个槽,逻辑位置通过 position id 来恢复。这样配合 RoPE 计算时,可以在读取时重新计算对应的旋转角度,保证位置信息不丢失。
3.3 环形缓存让旧缓存空间被复用
从内存管理角度,环形缓存本质是“预分配、复用”的思路。每次写入:
cache[head] = new_kv head = (head + 1) % W这个操作不仅写入新数据,还同时“释放”了最旧数据占用的逻辑空间,因为下一次写入会覆盖它。相比频繁append再pop(0)的队列实现,环形缓存没有元素搬移,也不需要动态扩容,对 GPU 显存分配器非常友好。
推理框架里常见做法是:为每条序列预分配一块(W, num_heads, head_dim)的连续显存。生成过程中,K/V 始终写在这块显存内部,指针循环回绕,不会触发新的显存分配。这一步对长连接、高并发场景特别重要,能有效避免显存碎片。
3.4 位置编码:环形缓存最容易踩的坑
使用环形缓存时,最容易出问题的是位置编码。
假设W=4,序列长度 6,第 6 个 token 需要关注第 3、4、5、6 个 token。写入环形缓存后,第 3 个 token 的 K/V 可能被覆盖到了第 5 个 token 的物理槽位。如果模型使用 RoPE,计算注意力时不能直接用物理槽位号5作为位置,因为第 3 个 token 的真实位置是2(从 0 开始计数)。所以要么在缓存中保存position,要么在写入时就把旋转位置相关的复数向量一起缓存。
这也是很多自研推理实现从“普通 KV cache”切到“环形 KV cache”后效果异常的原因:不是环错误,而是位置信息没有同步维护。
4. 完整实战:用 Python 实现环形 KV 缓存
4.1 项目结构与环境
为了把上面的原理落到代码里,我用一个纯 Python 示例演示环形 KV 缓存的基本工作方式。示例不依赖 PyTorch,只使用 Python 标准库,方便你在本地直接运行。
kv_ring_demo/ ├── ring_kv.py # 环形 KV 缓存实现 └── demo.py # 模拟 decode 生成流程环境要求:
- Python 3.8 或更高版本
- 无额外第三方依赖
如果你的项目里使用 numpy 或 PyTorch,把内部存储改为张量即可,逻辑是一样的。
4.2 环形 KV 缓存核心实现
先实现一个通用的环形 KV 缓冲区。每个写入项包含:
- key:当前 token 的 Key 向量,这里用
dim=4的随机向量模拟; - value:当前 token 的 Value 向量;
- position:当前 token 在原始序列中的真实位置 ID。
# ring_kv.py import random DIM = 4 def random_vec(): """生成一个随机向量,用来简化表示 Key/Value""" return [random.random() for _ in range(DIM)] class RingKVBuffer: def __init__(self, window_size: int): self.window_size = window_size self.keys = [None] * window_size self.values = [None] * window_size self.positions = [None] * window_size self.insert_count = [None] * window_size # 记录第几次写入,方便调试 self.head = 0 self.size = 0 self.num_writes = 0 def append(self, key, value, position: int): """写入新的 K/V,并覆盖最旧数据""" self.keys[self.head] = key self.values[self.head] = value self.positions[self.head] = position self.insert_count[self.head] = self.num_writes self.num_writes += 1 self.head = (self.head + 1) % self.window_size self.size = min(self.size + 1, self.window_size) def visible(self): """按时间顺序返回当前窗口中所有 K/V 及其真实 position""" start = self.head - self.size result = [] for i in range(self.size): idx = (start + i) % self.window_size result.append({ "key": self.keys[idx], "value": self.values[idx], "position": self.positions[idx], "insert_seq": self.insert_count[idx], "physical_slot": idx, }) return result def __len__(self): return self.sizeappend操作只做了两件事:把数据写入head指向的槽位,然后让head加 1 并对window_size取模。当缓冲区写满后,最旧的数据会被下一次写入自然覆盖。
4.3 模拟 decode 生成流程
接下来模拟一个文本生成过程。假设窗口大小为 4,我们连续“生成” 10 个 token。每个新 token 都有一个随机 Query 向量,它与缓存中所有可见 K 做点积,再与对应 V 做加权求和,得到一个粗略的注意力输出。
# demo.py import random import math from ring_kv import RingKVBuffer, random_vec, DIM random.seed(42) def dot(a, b): return sum(x * y for x, y in zip(a, b)) def calc_attention(query, cache: RingKVBuffer): """用缓存中的可见 K/V 计算简化注意力输出""" visible = cache.visible() if not visible: return None scores = [] for item in visible: score = dot(query, item["key"]) / math.sqrt(DIM) scores.append(score) # softmax max_score = max(scores) exp_scores = [math.exp(s - max_score) for s in scores] sum_exp = sum(exp_scores) # 加权求和 value output = [0.0] * DIM for item, exp_score in zip(visible, exp_scores): weight = exp_score / sum_exp for d in range(DIM): output[d] += weight * item["value"][d] return output def run_demo(): window_size = 4 cache = RingKVBuffer(window_size) print("=== 滑动窗口注意力 + 环形 KV 缓存 Demo ===") print(f"窗口大小 W = {window_size}\n") for step in range(10): # 生成当前 token 的 Key/Value/Position key = random_vec() value = random_vec() position = step # 真实位置 ID cache.append(key, value, position) # 当前 token 的 Query 也用随机向量模拟 query = random_vec() attn_output = calc_attention(query, cache) # 打印当前缓存状态 visible_items = cache.visible() visible_pos = [item["position"] for item in visible_items] print(f"step {step:2d} | 新 token pos={position} | " f"可见 position={visible_pos} | " f"物理槽位使用数={len(cache)}") print("\n最终缓存中保留的 position:", [item["position"] for item in cache.visible()]) print("最终缓存中保留的写入序号:", [item["insert_seq"] for item in cache.visible()]) if __name__ == "__main__": run_demo()运行方式:
python demo.py预期输出大致如下:
=== 滑动窗口注意力 + 环形 KV 缓存 Demo === 窗口大小 W = 4 step 0 | 新 token pos=0 | 可见 position=[0] | 物理槽位使用数=1 step 1 | 新 token pos=1 | 可见 position=[0, 1] | 物理槽位使用数=2 step 2 | 新 token pos=2 | 可见 position=[0, 1, 2] | 物理槽位使用数=3 step 3 | 新 token pos=3 | 可见 position=[0, 1, 2, 3] | 物理槽位使用数=4 step 4 | 新 token pos=4 | 可见 position=[1, 2, 3, 4] | 物理槽位使用数=4 step 5 | 新 token pos=5 | 可见 position=[2, 3, 4, 5] | 物理槽位使用数=4 step 6 | 新 token pos=6 | 可见 position=[3, 4, 5, 6] | 物理槽位使用数=4 step 7 | 新 token pos=7 | 可见 position=[4, 5, 6, 7] | 物理槽位使用数=4 step 8 | 新 token pos=8 | 可见 position=[5, 6, 7, 8] | 物理槽位使用数=4 step 9 | 新 token pos=9 | 可见 position=[6, 7, 8, 9] | 物理槽位使用数=4 最终缓存中保留的 position: [6, 7, 8, 9] 最终缓存中保留的写入序号: [6, 7, 8, 9]从输出可以看到两个关键现象:
- 前 4 步缓存逐步填满;
- 第 4 步之后,缓存中始终只保留最近 4 个 position 的 K/V,内存占用被锁定在固定大小。
这就是环形缓存在 decode 阶段的核心效果:随着窗口滑动,旧数据自动被覆盖,内存不增长。
4.4 与普通 KV cache 对比
为了更直观地看出差异,我再用一个普通追加式缓存做对比。普通 KV cache 就是往列表尾部不断追加:
class NaiveKVBuffer: def __init__(self): self.items = [] def append(self, key, value, position): self.items.append((key, value, position)) def visible(self): return [{"key": k, "value": v, "position": p} for k, v, p in self.items]在同样的 10 步生成中,普通 KV cache 的len会从 1 一直增加到 10,而环形缓存始终不超过 4。如果 tokens 数量是 10000,窗口是 1024,普通方式会保存 10000 份 K/V,环形缓存只会保存 1024 份。显存差距随序列长度放大。
5. 生产环境中的工程实现细节
5.1 主流推理框架怎么做
在实际的 LLM 推理框架中,环形 KV cache 通常不是用一个 Python 类管理,而是在显存分配阶段就预留好固定大小的张量块。
核心配置项通常包括:
# 伪代码,示意核心逻辑 kv_cache = torch.zeros( (window_size, num_heads, head_dim), dtype=torch.float16, device="cuda" ) position_cache = torch.zeros((window_size,), dtype=torch.long, device="cuda") def write_kv(kv_cache, key, value, position, step): slot = step % window_size kv_cache[slot] = key # 或把 key/value 分开存储 position_cache[slot] = position其中step是当前生成的第几个 token。通过step % window_size计算物理槽位,这就是环形缓存的本质。
在实现时,需要把 K 和 V 分开存储,方便与 Flash Attention 等高性能算子对接。每层模型需要维护一个独立的环形 KV 缓存,因为不同层的注意力权重不同。
5.2 与 Flash Attention / PagedAttention 的配合
现代推理框架通常用 Flash Attention 加速注意力计算。Flash Attention 本身支持传入自定义注意力掩码。当使用滑动窗口时,可以把窗口外部分置为负无穷,让 kernel 跳过这些位置。
但要注意:如果 KV cache 本身已经是环形缓冲,那传入 Flash Attention 的 K/V 张量可能需要在物理位置上做一次“重排”,把窗口内从旧到新的顺序恢复出来。否则计算索引会非常复杂。
还有一类方案是按页管理(类似 PagedAttention 的思路):把 KV cache 切分成固定大小的 page,由一个 block table 映射逻辑位置到物理位置。淘汰最旧 KV 时,只需要释放对应的 page,并把新 page 挂到末尾。这种方式相比纯环形更灵活,支持非连续显存,也更贴合 vLLM 等框架的做法。但核心思想仍然是“固定大小、循环复用”。
5.3 批量推理与多轮对话的边界
批量推理时,多条序列共享同一个模型权重,但各自的 KV cache 必须隔离。每条序列都要有自己的环形 KV cache,否则一个序列的覆盖会污染另一个序列。
多轮对话场景也要特别注意:滑动窗口不仅会丢弃很远的用户输入,还可能在极端情况下丢弃跨轮次的关键历史信息。此时工程上常见的做法是:
- 在滑动窗口层只缓存“局部上下文”;
- 在顶层或用少量全局 attention 层保存全局信息;
- 或者在多轮输入里显式使用摘要、记忆等机制,把关键信息压缩进窗口内。
6. 常见问题与排查思路
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 序列变长后显存仍然线性增长 | KV cache 使用的是普通追加式列表,没有复用旧槽位 | 改用预分配的环形 KV cache,确保写入走step % window_size |
| 生成内容从某个位置后开始错乱 | 环形缓存覆盖了仍在窗口内的 K/V,或位置索引没有同步维护 | 检查窗口计算边界;检查 position cache 是否正确写入 |
| 位置信息混乱,注意力分数异常 | 使用 RoPE 时只保存了物理槽位号,没有保存真实 position | 缓存中额外保存真实 position,或用 offset 计算相对位置 |
| 多 batch 推理时结果互相污染 | 多条序列共用同一个环形 KV cache | 每条序列独立维护自己的 KV cache |
| 切换环形缓存后性能反而下降 | 读出时需要把环形数据重排为连续顺序,增加了拷贝开销 | 评估重排成本;考虑使用 page-based 缓存代替纯环形 |
| 配置了滑动窗口但显存没有下降 | 只有注意力 mask 生效,缓存仍然保留全部历史 K/V | 对滑动窗口层裁剪 KV cache 保留范围 |
下面重点展开两个高频问题。
第一个是位置错乱。
很多自研实现会把slot = step % window_size当作 position 直接用于 RoPE 计算。这会在窗口滑动后出错。正确做法:保存position,并在计算 RoPE 时使用position,而不是slot。
第二个是窗口边界。
假设窗口大小为W,当前 token 位置是t,它可见的最早位置是t - W + 1。如果代码里写成t - W,就会多保留一个旧 token,虽然不会立即报错,但会让窗口语义偏移,且显存略高于预期。建议在单元测试中覆盖t < W和t >= W的边界。
7. 最佳实践与工程建议
围绕“滑动窗口注意力 + 环形 KV 缓存”落地,我给出几条工程建议。
第一,窗口大小必须和模型训练配置一致。
推理阶段的window_size不能随意调大或调小。如果训练时窗口是 1024,推理时突然改成 2048,模型并没有在这 2048 的范围内学习过长距离依赖,结果可能不增反减。同理,调小窗口会丢失训练时能看到的上下文,生成质量会下降。
第二,位置缓存和 KV 缓存必须一起维护。
凡是写入 K/V 的路径,都要同步写入真实 position。建议把 position 与 K/V 封装在同一个缓存对象里,避免不同代码路径漏写。
第三,预分配显存,避免动态分配。
环形缓存的优势之一就是显存固定。工程上应该在创建缓存时一次性分配(window_size, num_heads, head_dim)的连续空间,而不是每次append都触发内存分配。否则在高并发下容易产生显存碎片。
第四,优先考虑与算子库配合。
如果项目使用 FlashAttention、xFormers 等算子,优先看它们是否支持 sliding window 参数。如果支持,把环形的重排逻辑放在 kernel 外,用尽量少的拷贝把窗口内 K/V 恢复为连续张量。
第五,监控曲线要盯三件事。
- 长序列下 KV cache 显存是否保持水平;
- 每秒生成 token 数是否随序列长度明显下降;
- 多轮对话后恢复的 attention 分数是否出现异常峰值。
第六,不是所有层都需要滑动窗口。
Longformer 风格的模型中,常把一部分层设为全局注意力,一部分层设为滑动窗口注意力。对全局注意力层,KV cache 需要完整保存;对滑动窗口层,才使用环形缓存。混用时不要全局统一替换。
8. 总结与下一步
回到最开始那个问题:为什么滑动窗口注意力在 decode 时要使用环形缓存?
因为 decode 阶段是逐 token 增量生成的,默认的 KV cache 会随着序列长度线性膨胀,而滑动窗口注意力本身决定了每个 token 只能看到最近的 W 个 token。环形缓存用固定大小的存储空间、O(1) 的覆盖写入,恰好把“窗口滑出”的旧缓存立即复用,避免了显存浪费,也避免了列表头尾搬移的开销。
本文要点总结如下:
- 滑动窗口注意力把注意力的计算范围限制在最近 W 个 token;
- decode 阶段需要不断追加新 K/V,普通追加式缓存无法控制显存增长;
- 环形缓存通过
head = (head + 1) % window_size实现固定大小循环覆盖; - 实现环形 KV cache 时必须同步保存真实 position,否则 RoPE 等位置编码会错乱;
- 生产环境建议预分配显存、按序列隔离、与 Flash Attention 等算子配合使用。
下一步可以继续深入了解:
- RoPE 旋转位置编码的数学原理,以及它如何与环形缓存结合;
- Flash Attention 的滑动窗口掩码实现;
- PagedAttention 与环形缓存的设计差异;
- 多轮对话下的全局 token 与滑动窗口混合策略。
如果你正在自研推理引擎,或者想优化长文本服务的显存占用,可以从一个最小的环形 KV cache 开始改造,先在单层单序列上验证正确性,再逐步扩展到多层、多 batch。遇到“生成错乱”或“显存不降”的问题,优先检查窗口边界和 position 是否写对了。这两个点基本能覆盖 80% 的坑。