返回文章列表
记忆与上下文

PISA:用金字塔式 Top-K 将稀疏注意力降至对数线性复杂度

阅读约 3 分钟

导语

长上下文模型的一个核心难题,是标准自注意力需要计算所有查询与键之间的关联,序列长度增加时,计算和显存开销会快速上升。块稀疏注意力试图只保留少量重要的键块,但新的问题随之出现:如果仍要为每个查询评估全部键块,块的选择过程本身依旧接近二次复杂度。

来自 ByteDance Seed 的研究团队提出了 PISA,目标正是降低这一阶段的成本。论文题为《Block Sparse Attention with Log-Linear Complexity》,核心思路不是一次性在所有块中寻找 Top-K,而是借助多层次的候选筛选逐步缩小搜索范围。

核心方法:从粗粒度定位到细粒度选择

PISA 首先通过池化操作,把键表示组织成由粗到细的层级结构。较高层级包含更粗略、更少的键单元,较低层级则逐步恢复到更细的块。研究者据此构建出数量为 O(log N) 的层级。

对于每个查询,算法从最粗层开始进行候选选择。在当前层,它只对一个有界的候选集合计算 LogSumExp 分数,并保留其中的 Top-K 候选,再将这些候选向下一层展开。这个过程持续到最细层,最终得到用于稀疏注意力计算的键块。

这种金字塔式路由有两个直接效果:一是避免在每一层都扫描完整的键集合;二是把大规模搜索转化为多轮规模受控的局部选择。按照论文给出的分析,整体复杂度达到 O(N log N),其中 N 是序列长度。

工程实现与实验观察

为了让算法设计能够落地,论文还开发了面向硬件的 Triton 内核,覆盖训练和推理场景。实现将层级路由与 LogSumExp 打分进行融合,并避免显式生成完整的查询—键分数矩阵,这有助于减少中间结果带来的内存压力。

实验聚焦语言建模任务。根据论文摘要,PISA 与基线方法在常识推理等基准上的表现大致相当,同时在检索类任务上取得了更好的结果。这说明层级筛选并不只是追求计算量下降,在需要从长文本中定位相关信息的场景中,也可能带来实际收益。不过,摘要没有提供更细的模型配置、上下文长度、吞吐量或绝对分数,因此不宜仅凭现有材料判断其适用范围和实际加速幅度。

意义与待观察问题

PISA 的价值在于,它把长上下文注意力中的“如何选块”单独作为算法与系统协同优化的问题来处理。相比只减少最终注意力计算量,降低路由阶段的搜索成本,可能是块稀疏方法进一步扩展的重要前提。

同时,层级筛选也意味着模型需要依赖粗粒度表示提前判断相关区域。如果早期路由遗漏了关键信息,后续细化阶段可能无法恢复,因此候选数量、池化方式与不同任务之间的权衡仍值得进一步研究。总体而言,PISA 为长上下文模型提供了一条从全量匹配转向层级化检索的路径,其实际价值还需要更多规模、硬件和任务维度的公开评测来验证。

来源:Hugging Face Daily Papers

评论

正在确认登录状态……

正在加载评论……

相关文章

CCTest · Blog
LatentPort:让不同规模模型直接交接“记忆”,不再重放上下文
记忆与上下文
cctest.ai
记忆与上下文

LatentPort:让不同规模模型直接交接“记忆”,不再重放上下文

LatentPort 探索 Qwen3.5 4B 与 9B 之间的混合状态迁移:接收模型无需重新读取历史前缀,即可接管源模型处理后的部分运行状态。结果显示,仅迁移注意力 KV Cache 远远不够,加入 Gated DeltaNet 的持久状态后,续写损失明显下降。

阅读全文