滑动窗口注意力为何要用环形缓存?解码阶段显存优化实战
2026/8/31 3:00:09 网站建设 项目流程

最近在调一个长文本生成服务时,遇到一个非常典型的显存问题:模型启动后生成速度正常,一旦对话轮次变多、上下文变长,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+1t之间的 K/V 会被用到。比t-W+1更早的 K/V,从注意力计算角度已经彻底失去作用。

这带来一个重要推论:在 decode 阶段,只要保证缓存中保存最近 W 个 token 的 K/V,就足以计算结果。再早的 K/V 可以安全丢弃。

问题变成了:如何高效地“丢弃旧缓存、写入新缓存”?如果每步都通过移动数组元素来删除头部、追加尾部,即把后面的元素全部前移,会产生 O(W) 的搬运开销。而环形缓存通过覆盖写实现 O(1) 的淘汰和插入,正好匹配这个需求。

2.3 环形缓存解决的问题

综合来看,环形缓存解决的是三个问题:

  1. 内存固定:缓存数组预分配 W 个位置,不随序列长度增加而扩容,显存占用可控。
  2. 淘汰高效:写满后自动覆盖最旧数据,不需要移动元素,也没有额外释放/申请开销。
  3. 语义匹配:环形缓存的“滚动覆盖”天然对应滑动窗口的“窗口平移”。

因此,滑动窗口注意力在 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

这个操作不仅写入新数据,还同时“释放”了最旧数据占用的逻辑空间,因为下一次写入会覆盖它。相比频繁appendpop(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.size

append操作只做了两件事:把数据写入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]

从输出可以看到两个关键现象:

  1. 前 4 步缓存逐步填满;
  2. 第 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 < Wt >= 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% 的坑。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询