Research
MassAlloc Attention: Let Attention Allocate Its Own Compute
Overview Research area: Efficient attention mechanisms for long-context large language models, spanning attention-kernel design, sparse/adaptive compute allocation, and large-scale training and infere

- arXiv
- 2609.32712
- Published
- 2026-09-26
- Authors
- Jingze Shi, Zhangyang Peng, Xianduo Li, Yanlin Qi, Xiaotian Lin, Haoxian Chen, Liangdong Wang, Guang Liu, Yuyu Luo
AI summary
Overview
- Research area: Efficient attention mechanisms for long-context large language models, spanning attention-kernel design, sparse/adaptive compute allocation, and large-scale training and inference systems.
- Technical level: Advanced. The paper assumes familiarity with softmax attention, FlashAttention-style tiling, online softmax, and the backward pass of attention, though the central idea is describable in plain terms.
- Scope: The paper introduces MassAlloc Attention (MALA), a fused attention primitive that keeps full quadratic QK score discovery but uses normalized attention contribution to decide which post-score work to execute, and evaluates it across matched-work allocation controls, operator fidelity from 1K to 32K tokens, associative recall, 128K-token operator benchmarks on 8 H100 GPUs, scaling-law training from 0.6B to 14B parameters on 128 H100 GPUs, and model-level evaluation at 14B and 32B.
What This Paper Is About
Full softmax attention expends computation on every legal causal tile, even though much of the attention distribution carries negligible normalized mass. MALA keeps the ability to look at every causal interaction — it still computes all QK scores — but uses the attention distribution itself to skip the expensive computation that follows those scores.
The goal is to reduce training and inference cost, including the backward pass, without changing the structure of attention into a fixed window or a learned router, and without losing language-model quality or long-range retrieval.
Key Contributions
-
Distribution-conditioned compute allocation. The paper reframes sparse attention so that every causal interaction remains score-accessible, while the probability distribution realized by attention determines the post-score work that is executed. This contrasts with methods that restrict the support before exact attention and thereby skip QK computation for excluded regions.
-
A fused operator with paired online forward and offline backward rules under one tolerance. Forward uses the evolving online-softmax normalizer to test candidate tiles; backward reuses the finalized normalizer saved by forward to derive retained support nested within the forward support. A single normalized-mass tolerance (canonical value τ = 1) governs training forward and backward, inference prefill, and decoding.
-
Attention-state-only execution with a forward omitted-mass bound. The operators materialize neither an attention matrix nor a binary mask, retained-tile indices, or router state, and require only standard attention state — the running row maximum, the shifted normalization sum, and the output accumulator, with forward storing only the output and the final log-normalizer.
-
A broad empirical evaluation of modeling and efficiency consequences. The paper reports matched-work allocation controls against a per-instance reference-mass oracle, operator fidelity from 1K to 32K tokens, controlled associative recall, a 128K-token tensor-parallel operator benchmark, scaling-law training from 0.6B to 14B parameters, and knowledge, reasoning, and long-context retrieval evaluation of the resulting 14B and 32B models.
Main Findings
-
Matched-work allocation at 8K favors distribution-adaptive allocation. Under exactly matched total post-score work — a shared average of approximately 1,024 post-score key slots per query — MALA achieves mean omitted mass of 0.0188% versus 0.0182% for the per-instance reference-mass oracle, and mean relative output L2 error of 0.0174% versus 0.0164%. Position-only allocation yields 0.6012% mean omitted mass (P95 0.9923%) and 0.6231% mean output error (P95 4.237%), and layer/head/position allocation yields 0.1721% (P95 0.6247%) and 0.2817% (P95 0.9237%). Relative to static layer-head-position allocation, per-instance allocation reduces mean omitted mass by 9.5x and mean relative output L2 error by 17.2x.
-
One tolerance holds across context lengths in the operator-fidelity study. With the same tolerance from 1K to 32K tokens — a 32x increase in context — mean forward work rises from 470 to 1,066 key slots per query and mean backward work from 462 to 1,058; beyond 8K both remain between 1,010 and 1,066. Across all lengths, mean omitted probability mass is at most 0.0062% (P95 at most 0.032%), mean relative output L2 error is at most 0.021% (P95 at most 0.14%), and mean relative gradient errors are at most 0.38% for dQ, 0.35% for dK, and 0.17% for dV, with worst P95 errors of 0.90%, 0.66%, and 0.23%.
-
Associative recall tracks FullAttn as context grows. With 256 key-value pairs, sequence lengths from 1,024 to 8,192, and d_model from 64 to 512, and with fixed-budget sparse methods using a common 1,024-slot ceiling, MALA reaches 89.67% accuracy at sequence length 8,192 and d_model = 512 versus 89.97% for FullAttn. DSA reaches 52.61%, MoBA 47.25%, and NSA 22.61%.
-
Operator latency falls at 128K tokens while peak memory is retained. At 128K tokens on 8 H100 GPUs with tensor parallelism TP = 8, MALA reduces forward and backward latency during training by 2.2x and 3.0x and decoding latency during inference by 1.6x relative to FullAttn, while retaining FullAttn-level per-rank peak operator memory. Its decoding latency is within 3% of DSA and 1.6x faster than FullAttn, and it provides the lowest forward and backward latency among the compared methods.
-
Scaling-law training reduces FLOPs with comparable perplexity. Across five model sizes from 0.6B to 14B parameters on 128 H100 GPUs, MALA closely tracks FullAttn in perplexity. At 14B, perplexity differs from FullAttn by less than 0.001 after both pre-training and long-context training, while total training FLOPs are reduced by 2.5% during 4K pre-training and 23.1% during 32K long-context training, with complete QK score-discovery cost included in the accounting.
-
Model-level capability is preserved at 14B and 32B. At 14B, MALA scores 72.48 on knowledge and 64.66 on reasoning versus 72.32 and 64.46 for FullAttn; RULER at native 32K is 89.45 versus 89.42, and under YaRN extrapolation to 128K it is 65.75 versus 65.84. MoBA at 14B scores 70.01, 58.07, 86.85, and 60.17; DSA scores 70.11, 62.04, 87.59, and 63.05. At 32B, MALA reaches 75.75 knowledge, 76.10 reasoning, 92.71 RULER 32K, and 82.56 RULER 128K, versus 75.62, 75.67, 92.70, and 82.03 for FullAttn.
-
Complexity class is unchanged. MALA remains quadratic in sequence length because it computes QK scores for every legal causal tile; its savings are data-dependent constant-factor reductions in post-score arithmetic and memory traffic. As τ approaches 0, the skip conditions vanish and both passes recover FullAttn.
Methodology in Plain English
MALA divides attention into two stages. The first stage, score discovery, is untouched: for every legal causal tile, the kernel forms the QK score tile just as dense attention does. The second stage, post-score computation, is where allocation happens. In forward this stage covers the softmax update, loading V, and accumulating PV. In backward it covers probability reconstruction and the formation of dP, dS, dQ, dK, and dV.
The decision rule uses the attention distribution's own statistics. Uniform attention over a query row with L_q causally visible keys would give each key probability 1/L_q, so MALA scales that reference by a dimensionless tolerance τ to get a length-aware threshold τ/L_q. For each candidate tile, the kernel compares the tile's largest score contribution against the mass retained so far. In forward, the denominator is the evolving online-softmax normalizer, so the test is written as a comparison between the row maximum, the running row max, and the log of the running sum. If the resulting ratio falls below the threshold for every valid query row, the tile's post-score work is skipped and the running state is left unchanged; otherwise the standard online-softmax update proceeds. Causally masked entries have score minus infinity and do not constrain retention, and a row's first legal tile is always retained, so no attention row ends up empty. Traversal goes from the diagonal toward earlier keys, which initializes the running state from recent context and typically establishes a strong normalizer early — locality here is an execution prior, not a support assumption, since all earlier tiles are still scored and any distant tile with a large enough score is retained.
At the end of forward, the kernel overwrites the running sum buffer with the final log-normalizer, which is the only extra state passed to backward. Backward then applies the same tolerance using that finalized normalizer rather than replaying partial normalizers; "offline" in the paper means only that the normalizer is already finalized when a tile is tested, not that anything is calibrated ahead of time. Because the finalized normalizer is at least as large as any intermediate one, no key's contribution ratio can increase, so every tile forward skipped is also skipped in backward, and backward may additionally skip tiles forward retained before its normalizer was complete. The retained backward support is therefore nested inside the forward support without storing a mask. This nesting does not make backward an exact derivative of the forward operator, so gradient fidelity is measured empirically rather than asserted.
Work is measured as mean post-score key slots per query, excluding QK score discovery, where retaining a B_q x B_k tile contributes B_k post-score key slots to each of its B_q query rows. For cross-architecture comparisons with different value dimensions, the paper uses a separate post-score FLOP-equivalent accounting. The allocation tests are embedded in the ordinary tiled attention loop, and the same forward operator serves training and inference prefill. Under split-KV decoding, each split tests tiles against its own partial normalizer while L_q still counts the query's full causal context, which makes the test more conservative and may retain additional work; a shared tolerance therefore does not require split and unsplit execution to retain identical supports. MALA does not compress or evict the KV cache, so its decoding-memory scope is the operator's working set rather than cache capacity.
The evaluation compares MALA against FullAttn and against fixed-budget sparse methods. In the associative-recall study, fixed-budget methods use a 1,024-slot ceiling after causal clipping, including DSA despite its larger value dimension, while MALA uses the canonical tolerance. In the 128K operator benchmark, sparse configurations target approximately 1,024 reference-equivalent post-score slots per query, with MoBA using 1,024 raw slots and DSA using 256 because each slot contributes four times the reference PV work. Experiments use matched model scale, depth, hidden size, data, and optimization settings, with scaling-law and 32B continued training on 128 NVIDIA H100 GPUs and model-level evaluation and operator benchmarking on 8 H100 GPUs. The matched-work and operator-fidelity studies use the same 14B MALA checkpoint obtained after the 32K long-context stage of the scaling-law study.
Why This Matters
Impact on research. The paper argues for a different axis of sparsity than most prior work. Dynamic sparse attention typically reduces the search space before exact attention, which saves QK computation but risks discarding interactions; MALA keeps complete causal score discovery and instead allocates the post-score path, which is where a large share of the arithmetic and memory traffic in FlashAttention-style kernels sits. Because the rule is derived from the softmax normalizer the kernel already maintains, it produces nested forward and backward supports and a forward omitted-mass guarantee without a stored mask or a separately calibrated per-layer, per-head, or per-length threshold. The matched-work controls isolate why this matters: at identical total post-score work, static position-only and layer-head allocations perform substantially worse than a per-instance oracle, and MALA's online decisions recover nearly all of that gap.
Real-world applications:
- Long-document understanding over contexts from tens of thousands to 128K tokens, where the operator benchmark and RULER results are directly relevant.
- Multi-turn reasoning and agentic workloads, where autoregressive decoding dominates cost and the reported 1.6x decoding latency reduction applies.
- Repository-level code generation, a setting the introduction cites as a driver of long-context demand.
- Long-context retrieval and retrieval-augmented generation, where the associative-recall comparison specifically tests whether arbitrary long-range bindings survive the allocation decision.
Industry relevance. The reported savings are expressed in the units that matter for deployment: latency and peak memory for the complete attention operator including each method's own selection path, and total training FLOPs including complete score discovery and method-specific routing or indexing. MALA is described as compatible with tensor parallelism at TP = 8 and with split-KV decoding, and it retains FullAttn-level per-rank peak operator memory because allocation is fused into the attention loop. The paper also states that the 14B models and separately continued-trained 32B models retain comparable knowledge, reasoning, and long-context retrieval scores to FullAttn, which matters because a cheaper operator that degrades model quality would not be adopted. The code is open-sourced under the name flash-sparse-attention.
Future Directions
- Interactions with KV-cache management. MALA retains the full KV history for score discovery and value access, so the paper leaves pruning, compression, retrieval, and offloading of cached states as separate optimization opportunities that could compound with post-score allocation.
- Fusing allocation with support-sparse QK skipping. Because MALA keeps quadratic QK score discovery, a natural question is whether its normalized-contribution signal could guide or be combined with methods that skip QK computation, capturing both sources of savings.
- Behavior under distributed execution configurations. The paper notes that split-KV decoding makes the test more conservative because each split uses a smaller available normalizer, and that a shared tolerance does not force split and unsplit execution to retain identical supports. How allocation behaves across a wider range of parallelism and sharding schemes is left open.
- Whether the allocation signal can inform other parts of the model. The retained-support statistics are produced inside attention but are not used elsewhere; whether these signals could drive routing, caching, or capacity decisions at the layer or model level is not explored.
Target Audience
This paper is most useful to researchers and engineers working on efficient attention kernels and long-context language models — particularly those implementing FlashAttention-style fused kernels, designing sparse or adaptive attention methods, or running large-scale training and inference under tensor parallelism. It is also relevant to systems researchers interested in how runtime compute allocation can be derived from statistics a kernel already computes, and to practitioners evaluating whether a sparse attention approach preserves model quality at 14B and 32B scale before adopting it. Readers without background in softmax attention, online normalization, and the attention backward pass will find the methodology sections demanding, though the introduction, matched-work results, and model-level tables are accessible at a higher level.
Authors’ abstract
FullAttn often assigns negligible normalized mass to much of the causal score space, yet dense kernels execute the complete post-score path after forming each QK tile. We introduce MALA, a fused attention primitive that preserves score access to every legal causal interaction and uses normalized contribution to allocate post-score computation. Forward uses its evolving online-softmax normalizer, while backward reuses the finalized normalizer to derive nested retained support using only standard attention state. A common tolerance governs training and inference, allowing for adaptive retention of the work. MALA reduces low-contribution post-score computation. A matched-work study at 8K isolates the benefit of distribution-adaptive allocation: under exactly matched total post-score work, MALA approaches a per-instance reference-mass oracle, with mean omitted mass of 0.0188% versus 0.0182%. Across context lengths from 1K to 32K tokens, the same tolerance maintains low output and gradient errors relative to the reference. Across a broader controlled associative-recall comparison, MALA closely tracks FullAttn as context grows, reaching 89.67% accuracy at 8K compared with 89.97% for FullAttn. In an attention-operator benchmark at 128K tokens with tensor parallelism, MALA reduces forward and backward latency during training by 2.2x and 3.0x and decoding latency during inference by 1.6x relative to FullAttn. Across scaling-law training from 0.6B to 14B parameters, MALA closely tracks FullAttn in perplexity while reducing total training FLOPs. The resulting 14B models and 32B models from separate continued training achieve comparable knowledge, reasoning, and long-context retrieval scores to FullAttn. These results indicate that allocating post-score computation according to normalized attention contributions can retain the evaluated capabilities of FullAttn while reducing attention computation.