尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

滑动窗口注意力为何要用环形缓存?解码阶段显存优化实战

滑动窗口注意力为何要用环形缓存?解码阶段显存优化实战 最近在调一个长文本生成服务时遇到一个非常典型的显存问题模型启动后生成速度正常一旦对话轮次变多、上下文变长GPU 显存占用几乎线性上升最后直接 OOM。第一反应是 batch 开太大把 batch 调小后依然如此后来逐层打印显存占用才发现问题出在 KV cache 上——它不仅没有复用旧空间还在 decode 阶段不断追加新缓存把滑动窗口注意力的优势完全浪费了。这篇文章围绕一个核心问题展开为什么滑动窗口注意力Sliding Window AttentionSWA在 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是窗口大小。以W4为例第 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 大小约为2K 和 V× L × num_heads × head_dim × seq_len × 每个元素字节数当seq_len从 1024 涨到 8192KV cache 也涨 8 倍。如果是 7B 模型、32 层、4096 隐藏维度、又开了多 batch显存很容易被缓存占满。因此长文本推理场景中KV cache 的管理是性能优化的核心。2.2 滑动窗口注意力如何影响历史 KV滑动窗口注意力相当于给每个 token 规定了“有效视野”。窗口滑到第t个 token 时只有位置在t-W1到t之间的 K/V 会被用到。比t-W1更早的 K/V从注意力计算角度已经彻底失去作用。这带来一个重要推论在 decode 阶段只要保证缓存中保存最近 W 个 token 的 K/V就足以计算结果。再早的 K/V 可以安全丢弃。问题变成了如何高效地“丢弃旧缓存、写入新缓存”如果每步都通过移动数组元素来删除头部、追加尾部即把后面的元素全部前移会产生 O(W) 的搬运开销。而环形缓存通过覆盖写实现 O(1) 的淘汰和插入正好匹配这个需求。2.3 环形缓存解决的问题综合来看环形缓存解决的是三个问题内存固定缓存数组预分配 W 个位置不随序列长度增加而扩容显存占用可控。淘汰高效写满后自动覆盖最旧数据不需要移动元素也没有额外释放/申请开销。语义匹配环形缓存的“滚动覆盖”天然对应滑动窗口的“窗口平移”。因此滑动窗口注意力在 decode 阶段使用环形缓存不是技巧性优化而是逻辑上的必然选择窗口滑到哪缓存就跟到哪。3. 为什么 decode 阶段必须使用环形缓存3.1 窗口滑动的本质是“遗忘”很多初学者以为滑动窗口只是“注意力掩码不同”掩码里把窗口外位置置为负无穷即可。这在 prefill 阶段是对的但在 decode 阶段会带来浪费。考虑一个窗口大小W4的场景。当生成到第 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 位置编码环形缓存最容易踩的坑使用环形缓存时最容易出问题的是位置编码。假设W4序列长度 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 向量这里用dim4的随机向量模拟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(fstep {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 pos0 | 可见 position[0] | 物理槽位使用数1 step 1 | 新 token pos1 | 可见 position[0, 1] | 物理槽位使用数2 step 2 | 新 token pos2 | 可见 position[0, 1, 2] | 物理槽位使用数3 step 3 | 新 token pos3 | 可见 position[0, 1, 2, 3] | 物理槽位使用数4 step 4 | 新 token pos4 | 可见 position[1, 2, 3, 4] | 物理槽位使用数4 step 5 | 新 token pos5 | 可见 position[2, 3, 4, 5] | 物理槽位使用数4 step 6 | 新 token pos6 | 可见 position[3, 4, 5, 6] | 物理槽位使用数4 step 7 | 新 token pos7 | 可见 position[4, 5, 6, 7] | 物理槽位使用数4 step 8 | 新 token pos8 | 可见 position[5, 6, 7, 8] | 物理槽位使用数4 step 9 | 新 token pos9 | 可见 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), dtypetorch.float16, devicecuda ) position_cache torch.zeros((window_size,), dtypetorch.long, devicecuda) 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 个 tokendecode 阶段需要不断追加新 K/V普通追加式缓存无法控制显存增长环形缓存通过head (head 1) % window_size实现固定大小循环覆盖实现环形 KV cache 时必须同步保存真实 position否则 RoPE 等位置编码会错乱生产环境建议预分配显存、按序列隔离、与 Flash Attention 等算子配合使用。下一步可以继续深入了解RoPE 旋转位置编码的数学原理以及它如何与环形缓存结合Flash Attention 的滑动窗口掩码实现PagedAttention 与环形缓存的设计差异多轮对话下的全局 token 与滑动窗口混合策略。如果你正在自研推理引擎或者想优化长文本服务的显存占用可以从一个最小的环形 KV cache 开始改造先在单层单序列上验证正确性再逐步扩展到多层、多 batch。遇到“生成错乱”或“显存不降”的问题优先检查窗口边界和 position 是否写对了。这两个点基本能覆盖 80% 的坑。
返回列表