Skip to main content
Back to timeline
arXivSource publication:

CoWA splits causal-history coverage across KV heads, cutting 128K training forward and backward latency to roughly one-seventh of FullAttn

Synopsis

The work introduces CoWindow Attention (CoWA), in which all KV heads share near-diagonal and prefix-sink windows while complementary long-range windows partition the remaining history across heads, making full causal coverage a collective property of the head ensemble; in a window-matched 8K ablation CoWA reaches 89.73% versus 89.97% for FullAttn while duplicated long-range windows perform substantially worse, a 128K tensor-parallel operator benchmark reduces training forward and backward latency by 7.4x and 8.6x and decoding latency by 3.0x, and scaling-law training from 0.6B to 14B tracks FullAttn perplexity while cutting total training FLOPs by 28.5% in the 32K long-context stage.

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

Interpretation

CoWA reframes full causal coverage from 'every head sees the whole history' to 'the head ensemble covers the whole history': all KV heads share a near-diagonal window and a prefix-sink window, the remaining causal distances are partitioned into complementary long-range windows, and each KV head receives one, so each head attends sparsely to distant tokens while their union covers the full causal prefix. Prior fixed sparse patterns (sliding windows, global tokens, fixed blocks) or dynamic selection (retrieval scores, importance estimates, learned routers) either give every head the same finite horizon or introduce a separate selection problem; CoWA replaces content-dependent selection with a position-defined window rule that requires no learned router or indexer. The paper gives the union-of-visible-keys formulation and an attention-cost formulation showing that under the default equal-width allocation the union covers the full causal prefix; it also states that full coverage does not imply the same head-specific interactions or outputs as FullAttn, since each QO head still applies softmax over its own visible-key set.

A window-matched ablation separates window width from complementary allocation: holding per-head window widths fixed and replacing duplicated long-range windows with complementary ones across eight KV heads raises 8K recall monotonically from 21.32% (1 unique window, 12.5% coverage) to 32.92% (2), 52.32% (4) and 89.73% (CoWA, 100% coverage), against 89.97% for FullAttn and 5.34% for the near-only SWA reference. The ablation isolates the contribution of complementary long-range allocation itself rather than of window size or budget; duplicated long-range windows perform substantially worse at the same per-head widths. The ablation runs at 8K context with matched per-head window widths across 12.5%, 25%, 50% and 100% coverage; the controlled associative-recall study uses 256 key-value pairs, sequence lengths 1,024 to 8,192 and matched per-query token budgets of 1,024 to 1,920, where at 8K CoWA reaches 89.73%, DSA 53.71%, MoBA 50.12%, NSA 25.23%, and the remaining methods stay near 10%.

The position-defined rule also yields a regular execution structure: each head is described by a small set of contiguous windows, so excluded attention blocks are skipped directly rather than discovered through a dense mask or router; one rule governs training forward and backward, inference prefill and autoregressive decoding, and aligns with tensor-parallel KV-head sharding via global KV-head indexing. Compared with MoBA's block pooling, routing, TopK, rearrangement and merging, and DSA's lightning indexer, quantization, TopK and sparse MLA, CoWA derives its visible blocks directly from sequence position and global KV-head index, avoiding these auxiliary structures. An operator-level end-to-end benchmark at 128K tokens on 8 H100 GPUs with tensor parallelism, including sparse-pattern construction and data movement, shows training forward and backward latency reduced by 7.4x and 8.6x and decoding latency by 3.0x; per-rank peak operator memory matches FullAttn during training and is 8.4 MiB during decoding, below FullAttn and substantially below MoBA and DSA, while MoBA's 128K backward pass exceeds the per-rank memory limit.

Scaling-law training from 0.6B to 14B tracks FullAttn perplexity at lower total training FLOPs, with 3.1% and 28.5% reductions at 14B during 4K pre-training and 32K long-context training respectively; the resulting 14B models and separately continued-trained 32B models achieve comparable knowledge, reasoning and long-context retrieval scores to FullAttn. This connects operator-level savings to model-level capability retention: at 14B knowledge is 72.70 versus 72.32 and reasoning 64.87 versus 64.46, at 32B knowledge 76.07 versus 75.62 and reasoning 75.53 versus 75.67; native 32K RULER stays within 0.3 points of FullAttn at both scales, and under YaRN extrapolation to 128K CoWA reaches 66.60 versus 65.84 at 14B and 81.78 versus 82.03 at 32B. Scaling-law and continued-training runs use 128 H100 GPUs, while downstream evaluation and operator benchmarking use 8 H100 GPUs with tensor parallelism 8; each task reports means and standard deviations over five evaluation runs with seeds 0, 42, 233, 666 and 1234, and the paper notes that some differences are small relative to run-to-run variability while others are more pronounced.

Perspective

The result targets long-context language-model training and inference settings that use grouped-query attention with KV heads sharded across tensor-parallel ranks: the window rule is position-defined, shared across training forward and backward, inference prefill and autoregressive decoding, and keeps complementary assignments across tensor-parallel ranks via global KV-head indexing. It makes 'not every head needs the whole history' a trainable, executable default for teams that want to cut duplicated long-range access without introducing a router or indexer; equal-width allocation is used for load-balanced decoding, and equal-area allocation for training is explicitly outside the scope of this work. Operator-memory measurements concern the attention operator's working set rather than persistent KV-cache capacity, and window-aware offloading and prefetching are listed as future work.

What remains to watch: whether complementary long-range windows preserve capability on tasks and data distributions beyond those evaluated here, since the paper notes that full coverage does not imply the same head-specific interactions or outputs as FullAttn; some model-level task differences are small relative to run-to-run variability while others are more pronounced, for example at 14B and 128K the lower NIAH-MQ score alongside a higher RULER-CWE score, indicating task-specific trade-offs rather than uniform preservation; the 32B results come from a separate continued-training experiment and enter only the model-level evaluation, not the scaling-law analysis; and equal-area allocation for training, window-aware KV-cache offloading and prefetching, and gains at the level of persistent KV-cache capacity are all left to future work.

Sources