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

CoWA 让 KV 头分工覆盖全部因果历史,128K 训练前反向延迟降至 FullAttn 的约七分之一

核心概要

该工作提出 CoWindow Attention(CoWA),让所有 KV 头共享近邻窗口与前缀汇窗口、并把剩余远距离历史切成互补长程窗口分给不同头,使完整因果覆盖成为头集合的集体属性;在 8K 窗口匹配消融中 CoWA 达 89.73% 而 FullAttn 为 89.97%,重复长程窗口明显更差,128K 张量并行算子基准中训练前向与反向延迟分别降低 7.4 倍与 8.6 倍、解码延迟降低 3.0 倍,0.6B 至 14B 扩展训练困惑度贴近 FullAttn 且 32K 长上下文阶段总训练 FLOPs 减少 28.5%。

AI-generated editorial illustration: CoWindow Attention: Full Causal Coverage Is a Collective Property

深度剖析

CoWA 把完整因果覆盖从“每个头都看全历史”改为“头集合合起来看全历史”:所有 KV 头共享近对角窗口与前缀汇窗口,剩余因果距离被切成互补长程窗口,每个 KV 头分到一段,因此每个头对远距离 token 稀疏、但并集覆盖整个因果前缀。 此前固定稀疏模式(滑窗、全局 token、固定块)或动态选择(检索分数、重要性估计、学习路由器)要么给每个头同一有限视野,要么引入独立的选择问题;CoWA 用位置定义的窗口规则替代内容相关选择,不需要学习路由器或索引器。 论文给出可见键集合的并集公式与注意力代价公式,说明在默认等宽分配下并集覆盖完整因果前缀;同时指出完整覆盖不等于与 FullAttn 相同的头级交互或输出,因为每个 QO 头仍在自己的可见键集合上做 softmax。

窗口匹配消融把“窗口宽度”与“互补分配”分开:在每头窗口宽度相同的条件下,把 8 个 KV 头上的长程窗口从重复改为互补,8K 召回从 21.32%(1 个唯一窗口,12.5% 覆盖)单调升到 32.92%(2 个)、52.32%(4 个)和 89.73%(CoWA,100% 覆盖),而 FullAttn 为 89.97%,仅近邻的 SWA 为 5.34%。 该消融直接隔离了互补长程分配本身的贡献,而不是窗口大小或预算的贡献;重复长程窗口在相同每头宽度下表现明显更差。 消融在 8K 上下文、匹配每头窗口宽度下进行,覆盖 12.5%、25%、50%、100% 四档;受控联想召回使用 256 个键值对、序列长度 1,024 至 8,192、匹配每查询 token 预算 1,024 至 1,920,8K 时 CoWA 89.73%、DSA 53.71%、MoBA 50.12%、NSA 25.23%,其余方法接近 10%。

位置定义的规则同时给出可执行的规则结构:每个头由少量连续窗口描述,被排除的注意力块可直接跳过而不必用稠密掩码或路由器发现;同一规则贯穿训练前向与反向、推理预填充与自回归解码,并按全局 KV 头索引与张量并行 KV 头分片对齐。 与 MoBA 的块池化、路由、TopK、重排与合并,以及 DSA 的 lightning indexer、量化、TopK 与稀疏 MLA 相比,CoWA 的可见块直接由序列位置和全局 KV 头索引导出,避免这些辅助结构。 128K token、8 张 H100、张量并行的算子级端到端基准(含稀疏模式构造与数据搬运)显示:训练前向与反向延迟分别降低 7.4 倍与 8.6 倍,解码延迟降低 3.0 倍;每 rank 峰值算子内存在训练时与 FullAttn 持平,解码时为 8.4 MiB,低于 FullAttn 且明显低于 MoBA 与 DSA;MoBA 的 128K 反向超出每 rank 内存上限。

从 0.6B 到 14B 的扩展律训练中 CoWA 困惑度贴近 FullAttn 且总训练 FLOPs 更低:14B 时 4K 预训练减少 3.1%、32K 长上下文训练减少 28.5%;得到的 14B 模型与另行继续训练得到的 32B 模型在知识、推理与长上下文检索上取得与 FullAttn 可比的分数。 这把算子层面的节省连接到模型层面的能力保持:14B 知识 72.70 对 72.32、推理 64.87 对 64.46,32B 知识 76.07 对 75.62、推理 75.53 对 75.67;原生 32K RULER 两个规模均在 FullAttn 的 0.3 分以内,YaRN 外推到 128K 时 14B 为 66.60 对 65.84、32B 为 81.78 对 82.03。 扩展律与继续训练在 128 张 H100 上进行,下游评测与算子基准在 8 张 H100 上、张量并行度为 8;每个任务报告 5 次评测(种子 0、42、233、666、1234)的均值与标准差,并注明部分任务差异相对运行间波动较小、另一些更明显。

启示与展望

该结果面向采用分组查询注意力、按 KV 头做张量并行分片的长上下文语言模型训练与推理场景:窗口规则由位置定义,训练前向反向、推理预填充与自回归解码共用同一规则,并按全局 KV 头索引在张量并行 rank 间保持互补分配。它使“每个头不必都看全历史”成为可训练、可执行的默认选项,适用于希望在不引入路由器或索引器的情况下降低重复长程访问的团队;等宽分配用于负载均衡解码,训练用的等面积分配被明确列为超出本文范围。算子内存的测量针对注意力算子的工作集,而非持久化 KV 缓存容量;窗口感知的卸载与预取被列为未来工作。

仍待观察的是:互补长程窗口在本文评测之外的任务与数据分布上是否同样保持能力,因为论文指出完整覆盖不意味着与 FullAttn 相同的头级交互或输出;模型级评测中部分任务差异相对运行间波动较小、另一些更明显,例如 14B 在 128K 下 NIAH-MQ 较低而 RULER-CWE 较高,说明存在任务级取舍而非逐项一致保持;32B 结果来自单独的继续训练实验,只进入模型级评测而不在扩展律分析中;训练用的等面积分配、窗口感知的 KV 缓存卸载与预取,以及持久化 KV 缓存容量层面的收益,都还留作后续工作。

来源