MALA Lets Attention Decide Where Its Compute Should Go
Long-context attention is expensive for more than one reason. The QK score matrix is large, but dense kernels also perform the complete post-score path for every legal causal interaction. Many of those interactions ultimately receive negligible normalized attention mass, yet the kernel still spends compute and memory bandwidth processing them. MassAlloc Attention, or MALA, asks whether attention can use its own distribution to decide which of that work is worth keeping.
How MALA works
MALA does not begin by deleting causal connections. It preserves score access to the full legal causal region, then uses the contribution revealed by softmax normalization to allocate the later stages of attention more selectively.
- Forward allocation: the kernel uses the evolving online-softmax normalizer to estimate contribution while processing score tiles.
- Backward reuse: after the final normalizer is available, the backward pass derives nested retained supports from standard attention state rather than requiring a separate importance model.
- One tolerance for both modes: the same tolerance controls retention during training and inference, allowing the amount of work to adapt to the current attention distribution.
- Fused execution: the method is implemented as an attention primitive intended to convert fewer low-contribution post-score operations into actual kernel-level savings.
What the reported tests show
In a matched-work study at 8K context, MALA omitted an average of 0.0188% of attention mass, close to 0.0182% for a per-instance reference-mass oracle. Because total post-score work was matched, the comparison isolates the quality of distribution-adaptive allocation rather than simply rewarding a method that performs more computation.
Across context lengths from 1K to 32K, the same tolerance maintained low output and gradient errors relative to FullAttn. In a broader associative-recall comparison, MALA reached 89.67% accuracy at 8K, compared with 89.97% for FullAttn. These results suggest that reducing post-score work need not imply a large loss in the tested capabilities.
The systems results are more direct. On a 128K-token attention-operator benchmark using eight H100 GPUs with tensor parallelism, MALA delivered 2.2x lower forward latency and 3.0x lower backward latency during training, plus a 1.6x decoding speedup during inference, relative to FullAttn. Scaling experiments covered models from 0.6B to 14B parameters. For 14B training with a 32K context, MALA reduced total training FLOPs by 23.1% while maintaining comparable performance on the evaluated capabilities.
Why it matters
The central design choice is important: MALA is not a fixed sparse pattern and does not assume that the same tokens are unimportant for every query. It lets the current attention distribution determine how much post-score computation each instance receives. That makes the method potentially useful when long-context attention contains a large amount of low-impact history.
There are also practical conditions behind the reported gains. Adaptive retention must be efficiently fused into kernels; otherwise selection and control overhead can consume the saved work. The benefit may also vary with model architecture, task, and how concentrated the attention distribution is. MALA is therefore best understood as an adaptive execution strategy for attention, rather than a replacement for the attention formulation itself.
Source: Hugging Face Daily Papers
Comments
Checking sign-in status...
Loading comments...