开发者必读:NOSA-8B的CompressK模块实现与稀疏注意力本地性约束技巧
【免费下载链接】NOSA-8B项目地址: https://ai.gitcode.com/OpenBMB/NOSA-8B
NOSA是一种可训练的稀疏注意力机制,专为KV缓存卸载设计,具有明确的本地性约束,并搭配推理系统(NOSI)以实现其效率。它在1B/3B/8B规模的LLM上,相比FullAttn提升了解码吞吐量高达5.04倍,相比InfLLMv2提升1.92倍,相比ShadowKV提升1.83倍,同时改善了长上下文/长生成质量。
CompressK模块:高效KV压缩的核心实现
模块定义与核心参数
CompressK模块位于modeling_llama_long_infllmv2.py中,是NOSA-8B实现KV缓存优化的关键组件。其核心功能是通过分块平均池化实现键(K)张量的压缩,从而减少内存占用并提升推理速度。
class CompressK(torch.nn.Module): def __init__(self, head_num_k, head_dim, kernel_size, kernel_stride=16): super().__init__() self.kernel_size = kernel_size # 分块大小,默认32 self.head_num_k = head_num_k # 键注意力头数量 self.head_dim = head_dim # 每个头的维度 self.kernel_stride = kernel_stride # 分块步长,默认16前向传播流程解析
CompressK的前向传播包含三个关键步骤:
- 分块索引计算:通过
calc_chunks_with_stride函数根据序列长度、核大小和步长计算有效分块索引,实现带重叠的滑动窗口分块 - 关键向量提取:使用
index_select按计算出的索引提取关键分块 - 平均池化压缩:对每个分块执行均值池化,将
[l, block_size, h, d]形状的张量压缩为[l/stride, h, d]
def forward(self, k: torch.Tensor, cu_seqlens): # 计算分块元数据,支持步长 filtered_k_indices, cu_seqlens_compressed = calc_chunks_with_stride( cu_seqlens, self.kernel_size, self.kernel_stride ) # 提取过滤后的键向量 filtered_k = k.index_select(0, filtered_k_indices.view(-1)) # 分块并执行平均池化 filtered_k = filtered_k.view( filtered_k.shape[0] // self.kernel_size, self.kernel_size, self.head_num_k, self.head_dim ) compressed_k = filtered_k.mean(dim=1) return compressed_k, cu_seqlens_compressed在注意力机制中的集成
在LlamaAttention类初始化时,CompressK模块被实例化并与其他组件协同工作:
self.compress_k = CompressK( self.num_key_value_heads, self.head_dim, kernel_size=self.kernel_size, kernel_stride=self.kernel_stride )其中默认参数设置为kernel_size=32和kernel_stride=16,这种配置在保持信息损失最小化的同时实现了2倍的压缩比。
稀疏注意力的本地性约束实现
核心设计理念
NOSA的稀疏注意力机制通过显式本地性约束平衡效率与性能,主要体现在modeling_llama_long_infllmv2.py中的topk_sparse_attention函数实现。该机制结合了三种关键分块策略:
- 初始块(init_blocks):每个查询的初始分块数量,默认1
- 本地块(local_blocks):查询附近的本地分块数量,默认2
- 选择块(select_blocks):通过评分选择的全局分块
本地性约束的实现细节
本地性约束通过以下技术手段实现:
- 分块索引计算:
q_idx = cache_lens // block_size # 计算查询所在分块索引- 因果掩码应用:
j_idx = torch.arange(block_score_cis.shape[-1], device=block_score_cis.device).unsqueeze(0) ninf_mask = j_idx > q_idx.unsqueeze(1) # 构建本地性约束掩码 block_score_cis = block_score_cis.masked_fill(ninf_mask.unsqueeze(0), float('-inf'))- TopK选择与排序:
topk_idx = block_score_cis.topk(topk, dim=-1).indices.sort(-1).values topk_idx[topk_idx > q_idx[None, :, None]] = -1 # 过滤超出本地范围的分块参数配置与性能平衡
通过调整以下参数可以平衡模型性能与计算效率:
self.block_size = 64 # KV分块大小 self.window_size = 1024 # 本地窗口大小 self.local_blocks = self.window_size // self.block_size # 本地分块数 self.topk = 64 # 每查询选择的TopK分块数默认配置下,模型将注意力范围限制在1024 tokens的窗口内(16个64 token分块),同时通过TopK选择保留关键远程依赖,实现了本地性与全局信息的有效平衡。
实践应用与性能优化建议
模块使用场景
CompressK模块与稀疏注意力机制特别适合以下场景:
- 长文本处理任务(如文档摘要、代码分析)
- 资源受限环境下的LLM部署
- 需要高吞吐量的推理服务
性能调优关键参数
| 参数 | 作用 | 建议范围 |
|---|---|---|
| kernel_size | 分块大小 | 16-64 |
| kernel_stride | 分块步长 | 8-32 |
| block_size | 注意力分块大小 | 32-128 |
| topk | 稀疏选择分块数 | 32-128 |
部署注意事项
- 当处理特别长的序列时,建议增大
window_size以保留更多上下文信息 - 在GPU内存受限情况下,可减小
kernel_size或增大kernel_stride以提高压缩比 - 对于需要精确推理的任务,建议降低
topk值并增加local_blocks比例
总结
NOSA-8B通过CompressK模块实现的KV压缩与带本地性约束的稀疏注意力机制,为长上下文LLM推理提供了高效解决方案。这种设计不仅将解码吞吐量提升了1.92倍(相比InfLLMv2),还通过显式的本地性约束保持了长文本处理的质量。开发者可以通过调整分块大小、步长和TopK参数,在特定硬件环境和任务需求下实现最佳性能平衡。
完整实现细节可参考modeling_llama_long_infllmv2.py,更多技术背景请参见论文《NOSA: Native and Offloadable Sparse Attention》。
【免费下载链接】NOSA-8B项目地址: https://ai.gitcode.com/OpenBMB/NOSA-8B
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考