跳到主要内容
返回时间线
arXiv来源发表:

MALA 让注意力按归一化贡献自行分配后置计算,128K 训练前向/反向延迟降至 FullAttn 的 1/2.2 与 1/3.0

核心概要

该工作提出 MassAlloc Attention(MALA),在保留全部合法因果 QK 打分的前提下,用在线 softmax 归一化贡献决定哪些 tile 继续执行后置计算,在 8K 匹配工作量下平均遗漏质量 0.0188%(参考质量 oracle 为 0.0182%),1K–32K 上下文保持低输出与梯度误差,8K 联想召回 89.67%(FullAttn 89.97%),128K、8 卡张量并行下训练前向与反向延迟分别降低 2.2 倍与 3.0 倍、推理解码降低 1.6 倍,0.6B–14B 规模律训练困惑度贴近 FullAttn 且 14B 在 32K 长上下文阶段总训练 FLOPs 减少 23.1%,14B 与 32B 模型在知识、推理与长上下文检索上取得与

AI-generated editorial illustration: MassAlloc Attention: Let Attention Allocate Its Own Compute

深度剖析

MALA 把稀疏注意力重新表述为“分布条件下的计算分配”:每个合法因果交互仍完成 QK 打分,但后置计算(softmax 更新、加载、累加,以及反向的概率重建与 dQ/dK/dV 路径)只对归一化贡献超过阈值的 tile 执行。 与在打分前就裁剪支持集的固定窗口、块模式或动态路由方法不同,MALA 保留完整因果 QK 打分,把节省放在打分之后;与按层/头离线标定阈值的做法不同,它用同一个无量纲容差 τ 贯穿前向、反向、预填充与解码。 论文给出前向遗漏质量上界命题(Proposition 1),并在 256 条 8K 序列、58,720,256 次运行时分配决策上做匹配工作量对照实验。

在总后置工作量完全匹配(平均约 1,024 个后置 key slot/query)的 8K 对照中,MALA 仅凭在线状态就接近逐实例参考质量 oracle:平均遗漏质量 0.0188% 对 0.0182%,平均相对输出误差 0.0174% 对 0.0164%,而位置型与静态层-头-位置分配明显更差。 该对照把“按实例自适应分配”的价值与“在线决策带来的近似损失”分离开来,说明收益来自分布自适应而非单纯的工作量预算。 三种诊断对照均使用最终参考质量排序候选区域,只改变每个决策分配的工作量;数据无关的整数取整使静态分配的总 slot 数与 MALA 精确相等。

单一容差在 1K–32K 上下文保持算子保真:平均遗漏概率质量不超过 0.0062%(P95 不超过 0.032%),平均相对输出误差不超过 0.021%,平均相对梯度误差对 dV、dQ、dK 分别不超过 0.38%、0.35%、0.17%。 反向复用前向保存的最终 log-normalizer 推导嵌套保留支持,无需存储前向掩码,且反向支持嵌套于前向支持之内。 在 256 条混合知识、推理、检索序列的 1,024 至 32,768 前缀上,对同一查询、键、值与因果支持同时运行筛选算子与关闭筛选的参考算子,每个长度聚合 409,600 条 query-head-layer 与 81,920 条 KV-head-layer 测量。

在 128K token、8×H100、TP=8 的注意力算子基准中,MALA 相对 FullAttn 将训练前向与反向延迟分别降低 2.2 倍与 3.0 倍、推理解码延迟降低 1.6 倍,同时保持 FullAttn 级别的每卡峰值算子内存;0.6B–14B 规模律训练困惑度贴近 FullAttn,14B 在 4K 预训练与 32K 长上下文阶段总训练 FLOPs 分别减少 2.5% 与 23.1%,14B 与 32B 模型在知识、推理、原生 32K 与 YaRN 外推 128K 检索上取得相当成绩。 收益来自跳过低贡献后置执行,而 QK 复杂度仍为二次;与 MoBA、DSA 相比,MALA 不需要池化、路由、TopK 或索引器等额外结构,分配状态就是标准注意力状态。 延迟为张量并行组的同步墙钟时间,25 次预热后取 100 次迭代平均;模型级评测在 14B 比较四种注意力变体、在 32B 比较 FullAttn 与 MALA,每个任务报告 5 个随机种子的均值与标准差。

启示与展望

该结果面向需要长上下文训练与推理的注意力算子实现者与模型训练团队,适用场景是保留完整因果 QK 打分、以分块注意力循环执行前向、反向、预填充与自回归解码的设定;容差 τ 在训练与推理间共享,实际保留工作量由每层、每头、每条输入与每个序列长度上的注意力分布决定。它使后续工作可以在不引入路由或索引器的前提下,把后置计算按归一化贡献分配,并把节省转化为训练 FLOPs 与算子延迟的下降;解码侧的内存结论限于注意力算子的工作集,持久 KV 缓存容量不在其范围内。

读者仍需关注:MALA 保留完整因果 QK 打分,因此序列长度上的复杂度仍为二次,节省属于数据相关的常数因子;反向支持嵌套于前向支持之内,并不等于对前向算子的精确微分,梯度保真度由算子级测量给出;解码时若采用 split-KV,各分片以自身部分归一化器做判定,可能比非分片执行保留更多后置工作;论文明确把窗口感知的卸载与预取、以及降低 HBM 常驻 KV 存储列为未来工作,因此缓存容量与放置仍是独立问题;此外,14B 与 32B 的模型级结论来自特定评测套件与固定提示、样本数与解码设置,任务级差异(如 14B 在 128K 下 NIAH-MQ 较低而 RULER-CWE 较高)提示这是总体可比而非逐项一致。

来源