The question is when, not what

Most efficient attention work asks which keys and values a query should look at. Native Sparse Attention, Landmark Attention and S2-Attn all pick a subset of the past for every query and apply the same budget to every token. A new paper from Sakshi Choudhary, Aditya Chattopadhyay, Luca Zancato, Elvis Nunez, Matthew Trager, Wei Xia and Stefano Soatto at AWS Agentic AI asks a different question. Should this token look at the whole past at all?

Their example is a 100,000 token document. A sentence that continues 'therefore the quarterly revenue is' depends mostly on the paragraphs just above it. A phrase like 'as mentioned in Section 2' has to reach back tens of thousands of tokens. Standard attention pays the full quadratic price for both. The paper's bet is that the second kind of token is rare, and that a model can learn to recognise it.

The layer they propose is called L2A, for Learning To Attend. Every token first runs through sliding window attention over a 4K local window. The local output goes to a router, which is a single linear projection followed by a sigmoid. If the score is at or above 0.5 the token also runs exact global attention over the full context and the two outputs are summed. If not, the token keeps its local representation and moves on.

How they kept the router from collapsing

The training objective is ordinary next token prediction plus a penalty on the router scores, weighted by a coefficient lambda, so the model is pushed to invoke global attention sparingly. The routing decision is a hard threshold with zero gradient almost everywhere, so they use a straight-through estimator and backpropagate through the sigmoid instead.

The naive version of this collapses. The paper works through the gradient and shows why. When a token skips global attention, the global output is multiplied by zero, so the only gradient the router sees is the sparsity penalty, which says skip more. The router drifts to never invoking global attention and there is no signal to pull it back. Their fix is a coin flip. In each training step, with probability 0.1, they force global attention on for every token regardless of the router.

A few other choices matter for anyone trying to reproduce this. The router weights start at zero, which makes the sigmoid output 0.5 and triggers global attention for every token at initialisation, so training starts from something equivalent to the base model. They freeze the feed-forward layers and train only the token-mixing layers and layer norms, which they found works better than full fine-tuning.

What the numbers say

The base models are Qwen2.5 1.5B, Qwen2.5 7B and Qwen3 8B, all pretrained at 32K context, and the target is 128K. The comparison that matters is against continued long-context pretraining with full attention, which they call CLP and treat as the ceiling. Across HELMET, BABILong and MRCR, L2A lands within 1.5 to 3 percent of CLP while 75 to 80 percent of tokens skip global attention. Short-context benchmarks are unchanged within seed noise.

The sparse baselines do not fare as well. On Qwen2.5 7B, S2-Attn trains faster than L2A, at roughly 1725 versus 1384 tokens per GPU per second, but sits 10 points below CLP on the aggregate long-context score. The open-source NSA implementation they tested runs at about 565 tokens per GPU per second and is 20.5 points below CLP.

The ablation we found most convincing is the context-free baseline. They replaced the learned router with a Bernoulli coin that fires global attention with probability 0.1, 0.2 or 0.5, giving 50 to 90 percent sparsity with no dependence on the input. L2A at roughly 80 percent sparsity beats the best of these, the 50 percent sparsity variant, by nearly 30 percent on the 1.5B model. So the router is reading the local context and making a decision, rather than acting as a fancy dropout.

Where the speedup actually comes from

Algorithmic sparsity does not turn into wall-clock time on its own. GPUs like dense tiles, and conditional execution produces irregular memory access. The paper's second contribution is a Triton kernel built on the FlashAttention-2 tiling scheme. Active queries are gathered into a compact contiguous buffer, their original positions are kept in an index map, and causal masking uses the true positions rather than the compacted order. That lets the kernel skip key-value tiles that fall outside a query's causal range while keeping the same online log-sum-exp accumulation FA-2 uses for numerical stability.

On H200s, the standalone kernel runs 1.6 to 35 times faster than FA-2 in the forward pass and 1.08 to 45 times faster in the backward pass across sparsity levels, with roughly 10x and 8x at the 90 percent sparsity they see in practice. End to end, training throughput at 128K roughly doubles: 3059 versus 1518 tokens per GPU per second on Qwen2.5 1.5B, 1384 versus 643 on Qwen2.5 7B, and 410 versus 226 on Qwen3 8B.

The catch is the KV cache. A token that skips global attention still has to store its keys and values, because a later token might attend to it. So without further work L2A costs the same memory at decode time as the base model. Their answer is to measure per-layer sparsity after training and delete the global attention module from any layer where it fires less than 5 percent of the time. On the 1.5B model that removes 15 of 28 global modules, cuts KV cache by about half, and costs at most 1 percent on long-context scores. On the 7B model about 40 percent of layers can go.

What the router learned about tasks

The sparsity pattern is not uniform and the pattern is informative. Needle-style recall and MRCR come out very sparse, because the long-range dependency sits at a handful of positions. Retrieval-augmented generation and passage re-ranking land in the middle, since relevance comparison needs global access at many tokens. Many-shot in-context learning is the least sparse task, because comparing patterns across demonstrations means the local window is frequently insufficient.

Two loose ends. First, the sparsity level is set by lambda at training time and the paper admits it is hard to know what value to pick in advance. They show you can move the sigmoid threshold at test time to trade accuracy for speed per task, which is practical, though it means the reported numbers depend on a knob. Second, all of this is fine-tuning of pretrained models. Whether a router trained from scratch, or one combined with key-value sparsification as a second stage, holds the same gap to dense attention is untested here. The code is released under Apache 2.0, so someone could find out.

Sources

  1. Learning When to Attend: Conditional Memory Access for Long-Context LLMs (arXiv 2603.17484)