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 tt, the model processes input token xtx_t to predict xt+1x_{t+1}. Local and Full are two access modes of the same frozen backbone FϕF_\phi. The access sets are:

StF={1,…,t},StL={1,…,min⁡(s,t)}∪{max⁡(1,t−w+1),…,t}S^F_t = \{1, \ldots, t\}, \quad S^L_t = \{1, \ldots, \min(s, t)\} \cup \{\max(1, t-w+1), \ldots, t\}

where ss = initial positions (4) and ww = recent window width (2048).

Each step first executes Local, then computes a recall score:

(htL,ΔCtL)=FϕL(xt;Ct−1),qt=Rψ(ht−1,Eϕ(xt),htL)(h^L_t, \Delta C^L_t) = F^L_\phi(x_t; C_{t-1}), \quad q_t = R_\psi\left(h_{t-1}, E_\phi(x_t), h^L_t\right)

If qt>θq_t > \theta, ODA recomputes with Full; otherwise accepts Local. The default threshold is θ=0\theta = 0.

Learning to Predict Recall Benefit

The access gain measures the predictive benefit of Full over Local:

gt=ℓtL−ℓtF=log⁡ptF(xt+1)−log⁡ptL(xt+1)g_t = \ell^L_t - \ell^F_t = \log p^F_t(x_{t+1}) - \log p^L_t(x_{t+1})

With a Full-call penalty λ≥0\lambda \geq 0, the penalized one-step loss is:

ℓt(at)=ℓtL−at(gt−λ)\ell_t(a_t) = \ell^L_t - a_t(g_t - \lambda)

The head learns a regression score using a signed-log transform T(u)=sign(u)log⁡(1+∣u∣)T(u) = \text{sign}(u)\log(1 + |u|) and Huber regression with transition point 1:

L(ψ)=1∣E∣∑t∈EℓHuber,1(qt−yt)\mathcal{L}(\psi) = \frac{1}{|\mathcal{E}|}\sum_{t \in \mathcal{E}} \ell_{\text{Huber}, 1}(q_t - y_t)

Computational Savings

For a Full-call rate ρ\rho, ODA's average per-step savings relative to Full:

ΔF(n,ρ)=(1−ρ)FFull(n)−FLocal−FRecall\Delta F(n, \rho) = (1-\rho)F_{\text{Full}}(n) - F_{\text{Local}} - F_{\text{Recall}}

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):

Method4k16k64kAvg
Full-attn90.6381.9464.6179.98
StreamingLLM61.2919.239.1127.21
ODA90.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 tokensMethodFull (%)TPSSpeedupFLOPs saved (%)
516,788Full100.077.641.000×0.00
516,788Local0.0342.634.413×96.80
516,788ODA12.5205.712.650×84.31
516,788ODA25.0158.022.035×71.88
516,788ODA50.0107.291.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):

HELRULER16KFull (%)
✓——76.8483.20
—✓—49.8887.32
——✓81.6089.25
✓✓—81.9499.73
✓—✓67.9621.86
—✓✓81.5397.51
✓✓✓81.5776.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