NOSA-1B 核心组件揭秘:KV Cache 卸载如何带来 5.04× 解码吞吐提升? 📅 2026/8/17 19:31:04 NOSA-1B 核心组件揭秘KV Cache 卸载如何带来 5.04× 解码吞吐提升【免费下载链接】NOSA-1B项目地址: https://ai.gitcode.com/OpenBMB/NOSA-1B长文本生成慢显存不够用OpenBMB 开源的 NOSA-1B 给出了一个巧妙答案通过KV Cache 卸载加原生稀疏注意力把解码吞吐直接提升了5.04 倍对比 FullAttn。本文带你逐层拆解 NOSA-1B 的三大核心组件看懂这套原生可卸载稀疏注意力Native and Offloadable Sparse Attention的完整推理链路即使你是刚入门的大模型新手也能轻松理解其中的设计精髓。为什么长文本推理卡在 KV Cache大模型做长文本推理时每生成一个 token都要把历史 token 的 Key 和 Value 缓存下来这就是KV Cache。序列越长KV Cache 越大一段 128K 的对话KV Cache 可能占用几十 GB 显存。显存放不下就只能卸载到 CPU 内存可每次计算注意力都要把数据搬回 GPU——PCIe 带宽成了新的瓶颈推理速度断崖式下跌。传统的 KV Cache 卸载方案如 ShadowKV、InfLLMv2思路是全量卸载 按需取回但取回哪些数据往往靠启发式规则精度和效率都打了折扣。NOSA-1B 则换了一条路把哪些 KV 重要这件事直接让模型在训练时学出来。核心思想训练时学会挑选推理时只取精华NOSA-1B 的核心思想可以概括为三句话KV Cache 全部卸载到 CPUGPU 显存只留极少必需数据 训练时引入显式局部性约束让模型学会判断哪些 KV 块对后续生成真正重要⚡ 推理时只把最重要的少量 KV 块搬回显存参与注意力计算其余全部跳过。这套机制由可训练稀疏注意力 NOSA 与配套推理系统 NOSI 共同实现最终在 1B/3B/8B 三种规模模型上都验证了效果。接下来我们逐个揭开它的三大核心组件。核心组件一Key 压缩器如何生成 KV Cache 摘要第一个组件是Key 压缩器CompressK位于 modeling_llama_long_infllmv2.py。它的任务很直观把冗长的 Key 序列压缩成摘要方便后续快速筛选重要区域。它的工作原理是用calc_chunks_with_stride计算出需要做稀疏注意力的分块位置支持带步长的滑窗切分见 modeling_llama_long_infllmv2.py把 Key 按32 个 token 一组kernel_size32步长 16切块对每个块做均值池化得到压缩后的 Key 向量。压缩后的 Key 体积只有原来的约 1/32配合LlamaAttention中block_size64、window_size1024的分块配置见 modeling_llama_long_infllmv2.py模型可以在毫秒级完成全局粗筛。核心组件二CIS 重要性分数与 Triton 池化内核光有压缩 Key 还不够模型还需要知道哪些区域值得注意。第二个组件就是CIS压缩重要性分数一个可学习的打分模块为每个位置生成一个重要性分数分数越高的 KV 块越可能在后续生成中被用到。为了让这个打分足够快NOSA-1B 专门实现了一个Triton 池化内核完整代码在 cis_pooling.pynosa_mean_pool_kernel是一个 Triton JIT 内核按 (head, batch, window) 三维网格并行每个线程块读取固定大小的窗口分数并求均值支持变长序列通过cu_seqlens区分不同样本nosa_mean_pooling是对外的 Python 封装默认 kernel_size32、stride16与 Key 压缩器完全对齐。这套 Triton 实现相比朴素 PyTorch 版本避免了大量中间张量的显存开销是 CIS 打分能实时跟上解码节奏的关键。核心组件三两阶段 TopK 选择如何锁定关键 KV 块有了压缩 Key 和 CIS 分数接下来就是最精彩的部分两阶段 TopK 稀疏选择实现在 modeling_llama_long_infllmv2.py 的compressed_attention函数中。第一阶段按注意力分数粗选。用压缩后的 Key 和 Query 做一次轻量注意力得到每个块的粗分数block_score先选出前若干候选块qk_select。第二阶段按 CIS 分数精排。把 CIS 分数池化到同样粒度得到block_score_cis与第一阶段候选合并后最终选出topk64个最重要的 KV 块得到topk_idx。这个先粗后精的设计非常聪明粗选保证不漏掉注意力强的区域精排则把真正会复读的块挑出来两者互补。解码阶段如何只用 64 个 KV 块完成自回归生成选择完成后真正的注意力计算在sparse_forward中进行见 modeling_llama_long_infllmv2.py关键细节有三处CIS 参与 Value 加权scaled_v value_layer * cis重要性分数直接调制 Value让重要块在输出中占更大权重稀疏 FlashAttention只对topk_idx指向的 KV 块调用 varlen 版本的 FlashAttention 内核GPU 上真正参与计算的 token 数量大幅减少归一化修正用同结构但 Value 全为 1 的假注意力算出分母消除稀疏采样带来的偏差。每一步都在为少算、算准服务这就是解码吞吐飙升的底层原因。5.04× 解码吞吐提升是怎么来的数据对比官方在 1B/3B/8B 三种规模的模型上做了系统对比数据来自 README.md 的 Overview 部分对比方案解码吞吐提升特点FullAttn全注意力5.04×显存占用大长序列几乎不可用InfLLMv21.92×稀疏注意力但缺少可学习重要性ShadowKV1.83×KV Cache 卸载依赖启发式选择同时得益于训练时的显式局部性约束NOSA-1B 在长上下文/长生成场景下的质量也优于此前所有卸载方案——不仅快而且记得住。如何快速复现 NOSA-1B 并验证效果想亲自上手体验克隆仓库即可git clone https://gitcode.com/OpenBMB/NOSA-1B仓库内的关键文件一目了然modeling_llama_long_infllmv2.py稀疏注意力核心实现CompressK、compressed_attention、sparse_forward、SparseLlamaForCausalLMcis_pooling.pyCIS 分数的 Triton 池化内核modeling_minicpm.py基于 MiniCPM 架构的稀疏注意力变体config.json模型配置28 层、16 头、2 个 KV 头、longrope 外推到 32768 长度。权重文件pytorch_model.bin已随仓库发布配合 transformers 即可加载推理。总结NOSA-1B 用一套训练时学重要性 推理时只取精华的组合拳漂亮地解决了KV Cache 卸载场景下精度与速度的矛盾。Key 压缩器、CIS 打分内核、两阶段 TopK 选择这三大核心组件环环相扣共同铸就了5.04× 解码吞吐提升的亮眼成绩。如果你正在为长文本推理的显存与速度发愁不妨从读懂这几个核心组件开始把 NOSA-1B 的思路迁移到自己的项目里。【免费下载链接】NOSA-1B项目地址: https://ai.gitcode.com/OpenBMB/NOSA-1B创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考