Full text not available for this paper
Summary (Overview)
- On-Demand Attention (ODA) is a local-first decoding method that uses a lightweight "recall head" to decide when to invoke global attention during LLM decoding, keeping pretrained weights frozen.
- The key finding is that decoding states available after local computation in frozen pretrained models already contain signals predictive of the benefit of global attention over local attention.
- ODA recovers most of the performance lost under local attention while substantially reducing the frequency of global attention calls (e.g., ~41% on RULER16K for Qwen3-1.7B) and achieves up to 2.65× decoding throughput over full attention in controlled vLLM measurements.
- The method works across model scales and families (Qwen3, Qwen3.5, Gemma) and architectures (full-attention, hybrid attention), with only the recall head trained using modest data and compute.
Introduction and Theoretical Foundation
Motivation
Reasoning and agentic workloads demand efficient long-context inference. Full-attention decoding reads the growing history at every step, but the benefit of global access varies across prediction positions. The paper asks a key question: before computing global attention for the current step, can we estimate the additional predictive benefit it would provide over local computation and use this estimate to decide whether to access the complete history?
Theoretical Basis
The paper builds on the observation that local computation has a dual role: its output supports the current prediction AND, together with other available states, helps determine whether additional global computation is worthwhile. This contrasts with dynamic attention methods that jointly train the language model and its access policy—ODA instead keeps pretrained weights unchanged and learns to read out the benefit signal from existing model states.
Methodology
Local-First Decoding
At decoding position , the model processes input token to predict . Local and Full are two access modes of the same frozen backbone . The access sets are:
where = initial positions (4) and = recent window width (2048).
Each step first executes Local, then computes a recall score:
If , ODA recomputes with Full; otherwise accepts Local. The default threshold is .
Learning to Predict Recall Benefit
The access gain measures the predictive benefit of Full over Local:
With a Full-call penalty , the penalized one-step loss is:
The head learns a regression score using a signed-log transform and Huber regression with transition point 1:
Computational Savings
For a Full-call rate , ODA's average per-step savings relative to Full:
Efficient Execution
A vLLM runtime uses GPU-side conditional CUDA Graphs, a persistent ring KV workspace for Local, and fused KV commits across layers to reduce synchronization and kernel-launch overhead.
Empirical Validation / Results
Generation Quality (RULER, 4K–64K)
Table 1 (excerpt, Qwen3-1.7B):
| Method | 4k | 16k | 64k | Avg |
|---|---|---|---|---|
| Full-attn | 90.63 | 81.94 | 64.61 | 79.98 |
| StreamingLLM | 61.29 | 19.23 | 9.11 | 27.21 |
| ODA | 90.12 (54.69%) | 81.17 (41.64%) | 63.90 (69.63%) | 79.29 (58.43%) |
ODA preserves quality close to Full across all lengths, with score gaps below 1 point at 64K for all five models.
LongBench Results
On LongBench v1 (ten tasks), ODA's average score is within ~0.5 points of Full for all five models. For example, Qwen3-1.7B: Full = 39.50, ODA = 39.20 (80.12% Full-call rate).
State-Dependent Access Effectiveness
- At a 40% budget, the recall head captures 0.150 nats of net gain per reference position, 2.81× the expected gain of random selection.
- On five RULER16K tasks, ODA scores 88.96 at ~41% Full calls vs. Random's 32.79 at similar call rates.
Runtime Efficiency (Qwen3-1.7B, vLLM)
Table 20 (excerpt):
| Input tokens | Method | Full (%) | TPS | Speedup | FLOPs saved (%) |
|---|---|---|---|---|---|
| 516,788 | Full | 100.0 | 77.64 | 1.000× | 0.00 |
| 516,788 | Local | 0.0 | 342.63 | 4.413× | 96.80 |
| 516,788 | ODA | 12.5 | 205.71 | 2.650× | 84.31 |
| 516,788 | ODA | 25.0 | 158.02 | 2.035× | 71.88 |
| 516,788 | ODA | 50.0 | 107.29 | 1.382× | 47.01 |
At 512K input with 12.5% Full calls, ODA achieves 2.65× throughput and uses only 15.69% of Full's FLOPs.
Recall-Head Input Ablation
Table 3 (Qwen3-1.7B, RULER16K):
| H | E | L | RULER16K | Full (%) |
|---|---|---|---|---|
| ✓ | — | — | 76.84 | 83.20 |
| — | ✓ | — | 49.88 | 87.32 |
| — | — | ✓ | 81.60 | 89.25 |
| ✓ | ✓ | — | 81.94 | 99.73 |
| ✓ | — | ✓ | 67.96 | 21.86 |
| — | ✓ | ✓ | 81.53 | 97.51 |
| ✓ | ✓ | ✓ | 81.57 | 76.22 |
The complete input set (preceding hidden state, current embedding, current Local state) supports accepting more Local outputs while preserving quality.
Theoretical and Practical Implications
- State-guided computation: The paper provides empirical evidence that pretrained decoding states can serve dual purposes—both token prediction and decisions about accessing distant information—without additional training of the backbone.
- Architecture insights: Hybrid models (Qwen3.5 with Gated DeltaNet layers) require fewer Full calls than full-attention models, suggesting that recurrent states reduce the need for explicit global reads. Larger models within a family also achieve better quality-access trade-offs.
- Practical efficiency: The speedup depends critically on context length and recall frequency; at short contexts (4K), ODA is 13.1% slower than Full, but benefits grow substantially with context length.
- Access value concentration: Learned policies show large early gains followed by smaller increments, indicating that global access provides unequal benefits across positions—most quality recovery occurs at relatively low call rates.
Conclusion
ODA demonstrates that frozen pretrained decoding states contain signals predictive of the benefit of additional global attention. By training only a lightweight recall head, ODA recovers most of the generation quality lost under local attention with fewer global calls while retaining the complete history. The vLLM runtime with GPU-side conditional execution delivers practical decoding speedups in controlled long-context measurements.
Future directions identified by the authors include:
- Learning recall on policy-generated histories rather than reference histories
- Accounting for downstream generation quality and cumulative computational cost
- Improving selection accuracy (the offline Oracle still captures more benefit than the current recall head)
Related papers
- How Far Are We from Removing the Visual Encoder? Scaling Laws for Encoder-Free Multimodal Pretraining
Encoder-free multimodal LLMs match encoder-based performance at ~10^22 FLOPs, shifting compute-optimal allocation toward larger decoders and enabling viable encoder-free pretraining.
- Fast Learning Rate Transfer in Shallow Linear Networks at Growing Training Horizons
Fast learning-rate transfer holds at growing training horizons when T grows slower than sqrt(n), but requires nondegenerate first-order loss sensitivity to avoid spectral-dependent failures.
- Cost-free Spectral Estimation for Adaptive Newton--Schulz in Matrix Optimizers
Newton–Schulz iterations expose spectral moments via free scalar reductions, enabling spectrum-adaptive routine selection that cuts polar error up to 90x and improves LLM pretraining loss.