Skip to main content
Back to timeline
arXivSource publication:

MALA lets attention allocate its own post-score compute from normalized contribution, cutting 128K training forward/backward latency to 1/2.2 and 1/3.0 of FullAttn

Synopsis

The work introduces MassAlloc Attention (MALA), a fused attention primitive that keeps QK score discovery over every legal causal interaction and uses normalized online-softmax contribution to decide which tiles execute post-score computation; under exactly matched post-score work at 8K it reaches 0.0188% mean omitted mass versus 0.0182% for a per-instance reference-mass oracle, holds low output and gradient errors from 1K to 32K, reaches 89.67% associative-recall accuracy at 8K versus 89.97% for FullAttn, reduces training forward and backward latency by 2.2x and 3.0x and decoding latency by 1.6x at 128K with tensor parallelism on 8 GPUs, tracks FullAttn perplexity from 0.6B to 14B while cutting total training FLOPs by 23.

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

Interpretation

MALA reframes sparse attention as distribution-conditioned compute allocation: every legal causal interaction still undergoes QK score discovery, but post-score computation (softmax update, loading, accumulation in forward; probability reconstruction and the dQ/dK/dV paths in backward) runs only for tiles whose normalized contribution exceeds a threshold. Unlike fixed windows, block patterns, or dynamic routing that restrict support before scoring, MALA retains complete causal QK score discovery and places the savings after scoring; unlike offline per-layer or per-head threshold calibration, one dimensionless tolerance governs forward, backward, prefill, and decoding. The paper states a forward omitted-mass bound (Proposition 1) and runs a matched-work control study over 256 sequences of length 8,192, covering 58,720,256 runtime allocation decisions.

Under exactly matched total post-score work (a shared mean of about 1,024 post-score key slots per query) at 8K, MALA approaches a per-instance reference-mass oracle using only its online state: mean omitted mass 0.0188% versus 0.0182%, mean relative output error 0.0174% versus 0.0164%, while position-only and static layer-head-position allocations perform substantially worse. The control separates the value of instance-dependent allocation from the approximation introduced by deciding online, showing the gain comes from distribution adaptivity rather than from a work budget alone. All three diagnostic controls rank candidate regions by finalized reference mass and differ only in how much work each decision receives; data-independent integer rounding makes the static allocations match MALA's total slot count exactly.

A single tolerance preserves operator fidelity from 1K to 32K: mean omitted probability mass is at most 0.0062% (P95 at most 0.032%), mean relative output error at most 0.021%, and mean relative gradient errors at most 0.38% for dV, 0.35% for dQ, and 0.17% for dK. Backward reuses the finalized log-normalizer saved by forward to derive nested retained support without storing a forward mask, so backward support is nested within forward support. On 256 sequences drawn from a mixture of knowledge, reasoning, and retrieval data, prefixes of 1,024 to 32,768 tokens run the screened operator and a screening-disabled reference on identical queries, keys, values, and causal support, aggregating 409,600 query-head-layer and 81,920 KV-head-layer measurements per length.

In an attention-operator benchmark at 128K tokens on 8 H100 GPUs with tensor parallelism, MALA reduces training forward and backward latency by 2.2x and 3.0x and inference decoding latency by 1.6x relative to FullAttn while retaining FullAttn-level per-rank peak operator memory; scaling-law training from 0.6B to 14B tracks FullAttn perplexity while cutting total training FLOPs by 2.5% at 4K pre-training and 23.1% at 32K long-context training, and the resulting 14B and 32B models achieve comparable knowledge, reasoning, native 32K, and YaRN-extrapolated 128K retrieval scores. The savings come from skipping low-contribution post-score execution while QK complexity remains quadratic; unlike MoBA and DSA, MALA needs no pooling, routing, TopK, or indexer, because its allocation state is standard attention state. Latency is the synchronized wall-clock time of the tensor-parallel group, averaged over 100 iterations after 25 warm-ups; model-level evaluation compares four 14B attention variants and the FullAttn-MALA pair at 32B, reporting means and standard deviations over five seeds per task.

Perspective

The result targets implementers of attention operators and model-training teams working on long-context training and inference, in settings that keep complete causal QK score discovery and execute forward, backward, prefill, and autoregressive decoding inside a tiled attention loop; the tolerance is shared across training and inference, while realized retained work is set per layer, head, input, and sequence length by the attention distribution. It enables follow-up work to allocate post-score computation by normalized contribution without a router or indexer and to convert the savings into lower training FLOPs and operator latency; the decoding-memory claim concerns the attention operator's working set, not persistent KV-cache capacity.

Readers should still watch that MALA keeps complete causal QK score discovery, so complexity in sequence length remains quadratic and the savings are data-dependent constant factors; that backward support nested within forward support is not exact differentiation of the forward operator, with gradient fidelity reported at the operator level; that under split-KV decoding each split tests against its own partial normalizer and may retain more post-score work than unsplit execution; that the paper explicitly leaves window-aware offloading and prefetching to reduce HBM-resident KV storage as future work, so cache capacity and placement remain separate questions; and that the 14B and 32B model-level conclusions come from a specific benchmark suite with fixed prompts, shot counts, and decoding settings, where task-level differences such as the lower 14B NIAH-MQ and higher RULER-CWE at 128K indicate aggregate comparability rather than uniform per-task parity.

Sources