不同注意力头如何看到不同上下文?RedKnot SegPagedAttention分页注意力运行时设计揭秘
【免费下载链接】RedKnotEfficient Long-Context LLM Serving with Head-Aware KV Reuse and SegPagedAttention项目地址: https://gitcode.com/gh_mirrors/re/RedKnot
在长上下文大模型推理中,RedKnot 的 SegPagedAttention 分页注意力运行时把 KV 缓存拆分为「每个注意力头一张页表」,让 global 头读全量上下文、local 头只读 sink + 近期窗口,从而在不改变模型外部接口的情况下,把注意力头级的稀疏性转化为真实的显存与算力节省。本文面向新手,用 6 个部分讲清:为什么不同注意力头"看到"的上下文可以不同、分页存储如何组织、无 mask 的融合内核如何执行,以及实测收益有多少。
1. 为什么不同注意力头可以"看到"不同的上下文?
大模型的一个 Transformer 层里,注意力机制被拆成多个KV 头(KV head)。研究者们早已发现(DuoAttention 等工作):这些头并不是"人人平等"的——
- 有的头对前缀敏感:必须回看全部历史 token,才能保持输出质量,称为global(全局)头;
- 有的头对前缀鲁棒:只要看住开头几个 "sink" token 加上最近几百个 token,输出几乎不变,称为local(局部)头。
如果让所有头都用同一份"全量 KV 缓存 + 掩码(mask)"来执行,local 头其实做了大量无用功:不可见的 KV 照样占显存、照样被搬运。RedKnot 的做法很简单直接——既然 local 头根本不需要那些 token,那就别为它们分配存储空间。
这就引出了 SegPagedAttention 的核心思想:KV 缓存的存储粒度从「整个请求」细化到「层 × 注意力头 × 分段」。
2. RedKnot 的四种头分类策略
在 RedKnot 中,每个(层, KV头)组合都会被分配 4 种策略之一,定义在 head_config.py:
| 策略 | 可见范围 | 计算复杂度 | 角色 |
|---|---|---|---|
local | sink + 近期窗口 | O(L × W) | 长上下文的主力,省显存 |
global | 全量历史 KV | O(L²) | 保障前缀敏感信号 |
retrieval | top-p 重要 token | 稀疏检索 | 按需回看远距离关键 token |
dense | 全量 + RoPE 重对齐 | O(L²) | 浅层"安全网",保留细粒度信号 |
这些分类不是拍脑袋定的,而是离线剖析(profiling)得到的:RedKnot 提供 head_profiler.py 对真实模型逐头打分,输出 JSON 配置。仓库里也发布了多份现成的头分类配置,例如:
- Llama-3.3-70B 最优配置:llama-70B_optimal_g15_lf_ret.json(80 层 × 8 KV 头,约 15% global 头,local 窗口 8192,sink 4)
- DeepSeek-V4-Flash 发布配置:deepseek_v4_flash_0731_redknot.json
- Qwen3-32B / Mistral-7B 等配置:见 test/srt/redknot/head_class/
3. SegPagedAttention 分页存储:从"一张大表"到"每头一页表"
传统 PagedAttention 用一张全局页表映射「token → 物理页」。SegPagedAttention 则换成(层, 头, 分段) → 连续虚拟页的三维页表,核心实现在 segpaged.py:
SegmentPageTable(第 106-127 行):页表本体,记录每个头段的页号、每个头的执行策略(GLOBAL / LOCAL)及其分段列表;SegPagedKVCache(第 133-303 行):物理页池。local 头只分配sink + recent的页,global 头分配全量上下文的页。页与页之间通过虚拟→物理映射隔离,local 头的占用量与上下文总长无关;build_segpaged_cache(第 309-365 行):从稠密 KV 构建分头存储的入口,按头策略裁剪 local 头只保留可见 token。
一个形象的类比:传统缓存像图书馆的一整面书架,所有读者都面对全部藏书;SegPagedAttention 则像给每位读者配了自己的迷你书架——local 头这位"读者"的书架上永远只有开头几本和最近几本,无论图书馆涨到多大。
4. 无 mask 的融合内核:如何把稀疏性变成真实加速?
有了分头页表还不够——如果内核仍然按"最长的那个头"的长度去算,再靠 mask 屏蔽,节省不了算力。RedKnot 的segpaged_attention(第 428-547 行)执行的是论文算法 2:
- 逐头查页表,取出该头真实的 KV 长度;
- 把所有 (查询头, KV头) 对打包进一次FA-3 varlen(变长)调用,各头各算各的长度,全程无 mask;
- Hopper GPU 上走融合内核,其他环境自动回退到等价的 PyTorch 参考实现,保证可验证。
配套的verify_against_dense(第 592-672 行)会用稠密 + mask 基线做数值对照,论文要求余弦相似度 > 0.99998——也就是"换掉存储布局,结果不变"。
5. 实测收益:提速与 KV 节省到底有多少?
单张 H200 上的实测记录见 SEGPAGED_RESULTS.md,结论(注意力算子级):
| 场景 | 阶段 | 注意力算子提速 | KV 节省 | 数值等价 |
|---|---|---|---|---|
| H2O / Heavy-Hitter | decode | 4.8x ~ 14.9x(随上下文增长) | 90–96% | cos = 1.0 |
| DuoAttention | prefill | 1.8x ~ 8.3x | 72–75% | cos ≈ 1.0 |
几个值得注意的细节:
- 上下文越长,dense 基线要读的 KV 越多,差距越大(8K 时 4.78x → 32K 时 14.9x);
- 页表等布局开销必须在 KV 写入时一次性完成并复用,每次 attention 现算会吃掉收益(文中记录了 0.70x 的反例);
- KV 容量节省直接转化为并发能力:论文口径下 head-aware 调度可带来 4.7–7.8x 并发容量。
6. 如何动手体验 SegPagedAttention 后端?
SegPagedAttention 已注册为独立的 SGLang 注意力后端segpaged,入口在 segpaged_backend.py,启动方式:
python -m sglang.launch_server \ --attention-backend segpaged \ --segpaged-head-config-path /path/to/head_config.json \ --segpaged-page-size 64如果没有自己的头配置,也可以克隆仓库后直接跑演示脚本(完整说明见 examples/redknot/README.md):
git clone https://gitcode.com/gh_mirrors/re/RedKnot cd RedKnot CUDA_VISIBLE_DEVICES=0 python examples/redknot/segpaged_redknot_demo.py --mode smoke该 demo 会对比三条路径——稠密 + mask、PyTorch SDPA + mask、SegPaged 融合内核——并输出延迟、余弦等价性和 KV token 节省,适合新手建立直观感受。
小结
SegPagedAttention 的"分页"二字,本质是把注意力头级的异构性从"计算时的掩码"下沉为"存储时的物理隔离":global 头存全量页、local 头只存 sink + 近期页、检索头按需取页。配合无 mask 的 varlen 融合内核,同样的模型输出,换来 72–96% 的 KV 节省与数倍注意力加速。对长上下文 RAG、多轮对话这类业务,这正是 TTFT 与并发容量提升的底层来源。想深入了解论文级细节,可参考 ROADMAP.md 中的分阶段实现路线。
【免费下载链接】RedKnotEfficient Long-Context LLM Serving with Head-Aware KV Reuse and SegPagedAttention项目地址: https://gitcode.com/gh_mirrors/re/RedKnot
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考