公司动态

开发者必读:NOSA-8B的CompressK模块实现与稀疏注意力本地性约束技巧

📅 2026/8/14 19:54:59
开发者必读:NOSA-8B的CompressK模块实现与稀疏注意力本地性约束技巧
开发者必读NOSA-8B的CompressK模块实现与稀疏注意力本地性约束技巧【免费下载链接】NOSA-8B项目地址: https://ai.gitcode.com/OpenBMB/NOSA-8BNOSA是一种可训练的稀疏注意力机制专为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_stride16): 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(dim1) return compressed_k, cu_seqlens_compressed在注意力机制中的集成在LlamaAttention类初始化时CompressK模块被实例化并与其他组件协同工作self.compress_k CompressK( self.num_key_value_heads, self.head_dim, kernel_sizeself.kernel_size, kernel_strideself.kernel_stride )其中默认参数设置为kernel_size32和kernel_stride16这种配置在保持信息损失最小化的同时实现了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], deviceblock_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-64kernel_stride分块步长8-32block_size注意力分块大小32-128topk稀疏选择分块数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),仅供参考