MALA:让注意力权重决定哪些计算值得继续
长上下文注意力的计算瓶颈,不仅来自 QK 打分本身,也来自打分之后对大量低贡献位置执行的乘加、归一化和反向传播。MassAlloc Attention(MALA)提出了一个直接的思路:既然注意力分布已经反映了不同区域的重要性,就让注意力自身决定后续计算资源应该投向哪里。
核心方法
传统 FullAttn 会为所有合法的因果交互完成完整的后处理,即使某些位置最终获得的归一化注意力质量几乎可以忽略。MALA 不在 QK 阶段提前删除连接,而是保留对整个合法因果范围的打分访问;随后利用 softmax 的归一化贡献,筛选并保留更值得继续处理的区域。
- 前向阶段:利用在线 softmax 不断更新的归一化因子,动态判断各个区域的贡献,并分配后续计算。
- 反向阶段:复用已经确定的归一化因子,从标准注意力状态中推导出嵌套的保留支持集,避免引入额外的模型状态。
- 统一容差:训练和推理使用共同的误差容忍标准,使计算量可以随当前注意力分布自适应变化。
- 算子级优化:MALA 被实现为融合注意力原语,目标不是改变注意力形式,而是减少低贡献区域的 post-score 工作。
实验结果
在 8K 上下文的匹配工作量实验中,MALA 的平均遗漏注意力质量为 0.0188%,接近按实例使用真实参考质量的理想分配器(0.0182%)。这说明,在总后处理工作量完全匹配的条件下,MALA 的分配策略已经接近一个了解最终分布的参考方案。
从 1K 延伸到 32K 上下文时,同一容差仍能维持较低的输出误差和梯度误差。在更广泛的关联召回测试中,8K 上下文下 MALA 的准确率为 89.67%,FullAttn 为 89.97%,两者差距有限。
性能方面,在 128K token、8 张 H100、TP=8 的注意力算子基准中,MALA 相比 FullAttn 的训练前向、训练反向和推理解码加速分别为 2.2 倍、3.0 倍和 1.6 倍。规模化训练实验覆盖 0.6B 至 14B 参数模型;在 14B、32K 上下文训练中,MALA 将总训练 FLOPs 减少 23.1%,同时保持与 FullAttn 相当的已评估能力表现。
意义与限制
MALA 的价值在于,它没有简单地把注意力变成固定模式的稀疏连接,而是根据每个实例、每个阶段的注意力质量动态分配计算。这种设计尤其适合长上下文场景:当大量历史位置对当前输出贡献很小,节省的就不只是存储访问,还包括打分后的核心计算和训练反向开销。
不过,MALA 的收益依赖注意力分布是否足够集中,也依赖融合内核能否把动态筛选转化为硬件上的真实效率。因此,它更像是对 FullAttn 的自适应执行层改造,而非一种完全不同的注意力架构。其后续价值,将取决于不同模型、任务和上下文长度下的分布稳定性,以及在真实服务系统中的调度开销。
评论
正在确认登录状态……
正在加载评论……