Skip to content
AI.info

Research

Token Sparse Attention: Efficient Long-Context Inference with Interleaved Token Selection

Token Sparse Attention: Efficient Long-Context Inference with Interleaved Token Selection Authors: Dongwon Jo, Beomseok Kang, Jiwon Song, Jae-Joon Kim (Department of Electrical and Computer Engineerin

Token Sparse Attention: Efficient Long-Context Inference with Interleaved Token Selection
arXiv
2602.03216
Published
2026-02-03
Authors
Dongwon Jo, Beomseok Kang, Jiwon Song, Jae-Joon Kim

AI summary

Token Sparse Attention: Efficient Long-Context Inference with Interleaved Token Selection

Authors: Dongwon Jo, Beomseok Kang, Jiwon Song, Jae-Joon Kim (Department of Electrical and Computer Engineering, Seoul National University) arXiv: 2602.03216v3 [cs.CL], 29 May 2026 | License: CC BY 4.0 | Code: https://github.com/dongwonjo/Token-Sparse-Attention

Overview

Research area: Natural Language Processing — efficient long-context inference for large language models, specifically token-level sparsification of the attention mechanism during prefill.

Technical level: Intermediate. The paper assumes familiarity with transformer attention, the prefill/decode split, FlashAttention-style kernels, and the distinction between block-level and token-level sparse attention. The core idea, however, is intuitive enough to follow without kernel-level expertise.

Scope: The paper proposes and evaluates a reversible, per-head token-selection mechanism that compresses Q, K, and V to a reduced token set before attention and scatters the output back to the full sequence length afterward.

What This Paper Is About

Attention cost grows quadratically with context length during prefill, which is the central bottleneck for long-context LLM inference. Existing accelerations either sparsify the attention map at the block level (so unimportant tokens survive whenever they share a block with salient ones) or permanently evict tokens at early layers (so evicted tokens can never be reconsidered). This paper introduces Token Sparse Attention, which selects a per-head subset of tokens for each attention computation and then interleaves the output back into the original sequence dimension, letting token importance be re-evaluated across layers and heads rather than decided once and irreversibly.

Key Contributions

  1. A reversible "compress and then decompress" token-sparsification mechanism. Each attention head selects its own token index set from the full sequence, gathers the corresponding rows of Q, K, and V into compact tensors, performs dense attention in the reduced space, and scatters the result back into a zero-initialized tensor of the original sequence shape so the residual connection preserves unselected tokens.

  2. Dynamic Token Coverage, an inference-time sparsity-budget policy. Rather than fixing a retention ratio, the method estimates per-head token importance from a small set of recent queries, aggregates head scores into a layer-level distribution, sorts tokens by ascending importance, and keeps the minimal number of tokens whose cumulative importance exceeds a coverage threshold τ.

  3. A drift-based criterion for choosing which layers to sparsify. The authors define Inter-Layer Representation Drift, the relative L2 change between a token's input and output hidden states, and show empirically that layers with lower drift tolerate sparsification better. A normalized drift rank with δ = 0.5 selects the sparse layers, computed once per model as preprocessing.

  4. Demonstration that the method is complementary rather than competitive. Because it performs dense attention in the compressed space, it composes with FlashAttention and with block-sparse kernels such as Minference and FlexPrefill without kernel modification, producing heterogeneous granularity (block-sparse plus token-sparse) as a new design point.

Main Findings

  • Headline result: Up to 3.23x attention speedup at 128K context with less than 1% accuracy degradation (abstract). Gains grow with context length and are substantially larger at 128K and 256K, where attention dominates latency.

  • Complementary accuracy preservation on RULER: For LLaMA-3.1-8B-Instruct, FlashAttention averages 87.01% and reaches 87.02% with Token Sparse Attention; Minference goes from 86.49% to 86.05%; FlexPrefill remains unchanged at 87.27% with and without the addition. All three baselines gain speedup at 128K: FlashAttention from 1.00x to 1.36x, Minference from 1.12x to 1.38x, FlexPrefill from 2.44x to 2.76x.

  • Heterogeneous granularity beats tuning alone: Combining FlexPrefill (block-sparse) with Token Sparse Attention (token-sparse) reaches 87.3% accuracy at 2.8x speedup over FlashAttention, whereas standard FlexPrefill reaches the same 87.3% accuracy at 2.4x speedup.

  • Mistral-Nemo-12B-Instruct results: RULER averages are 67.60 (FlashAttention) to 67.37 (with Token Sparse, 1.22x), 66.42 to 66.46 (Minference, 1.28x), and 68.16 to 67.91 (FlexPrefill, 1.33x). Deviation introduced by the method stays within 0.5% of the baseline.

  • InfiniteBench averages are essentially flat: LLaMA-3.1-8B-Instruct goes 50.86 to 50.88 (FlashAttention), 50.16 to 49.70 (Minference), and 49.53 to 49.23 (FlexPrefill). Mistral goes 22.93 to 22.80, 19.77 to 19.41, and 24.80 to 24.08.

  • LongBench averages (Appendix A.2, LLaMA-3.1-8B-Instruct): FlashAttention 49.28 to 48.96, Minference 47.89 to 47.61, FlexPrefill 47.48 to 47.39, described as minimal accuracy changes across all baselines.

  • Sparsity rises with context length (Table 3): At τ = 0.005, average attention-map sparsity climbs from 17.00% at 4K to 54.44% at 128K. At τ = 0.010 it climbs from 28.02% at 4K to 67.36% at 128K.

  • Overhead is bounded: At 128K, the extra cost of token scoring/indexing plus QKV compression and output decompression accounts for less than 11% of total attention latency across all layers, even at the highest sparsity setting.

  • Dynamic beats fixed sparsity: At comparable speedup (1.36x dynamic vs. 1.32x fixed; 1.51x dynamic vs. 1.57x fixed), dynamic coverage at τ = 0.005 achieves 87.02% RULER average versus 86.91% for fixed s = 0.3, and dynamic τ = 0.010 achieves 86.84% versus 85.43% for fixed s = 0.5 at higher measured sparsity (74.95% fixed vs. 67.36% dynamic).

  • Beats token eviction under matched speedup (Table 5, LLaMA-3.1-8B-Instruct): At roughly matched speedups (1.49x PyramidInfer, 1.50x FastKV, 1.53x GemFilter, 1.51x Ours), Token Sparse Attention attains the highest average RULER accuracy at 86.84%, versus 78.49% (PyramidInfer), 85.12% (GemFilter), and 85.64% (FastKV), against the 87.01% FlashAttention reference.

  • Token importance is genuinely unstable (motivation): Tracking the top 1% tokens per layer in LLaMA-3.1-8B-Instruct, overlap between adjacent layers is moderate but decays rapidly with layer distance; at layer 18, different heads rank tokens differently, motivating per-head rather than shared token sets.

  • Scoring design matters (Table 7): Replacing recent-query scoring with random query selection drops RULER average from 87.02% to 84.95%; query-only pooling lands at 86.43%, below the recent-query design.

  • Drift predicts robustness: Across 200 random 3-layer sparsification runs at token coverage 0.99 on 4K-length RULER, layer subsets with lower mean normalized drift tend to yield higher accuracy. The drift pattern is described as consistent across tasks and context lengths tested; the appendix discussion of Low/Mid/High drift groups is truncated in the provided content.

Methodology in Plain English

The method intervenes at two points in each attention layer.

Step one — pick and compress. For every attention head, the model needs a way to decide which tokens matter. It takes a small set of the most recent queries and runs them against all keys to build a cheap proxy attention map. Summing those attention weights down the query dimension gives each token a per-head importance score. The scores are aggregated across heads and normalized into a layer-level distribution, then sorted from least to most important. Instead of asking "how many tokens should I keep?", the method asks the reverse: starting from the least important token, how few tokens must be dropped before their combined importance mass reaches a threshold τ? Those tokens define a budget, and each head then independently keeps its own top-scoring tokens within that budget. The selected rows are gathered out of the original Q, K, and V tensors into compact, dense, contiguous tensors.

Step two — attend and scatter back. Because the compact tensors are dense and contiguous, they can be fed directly into unmodified FlashAttention or an existing sparse kernel; no kernel surgery is required. Working in the reduced space cuts the quadratic attention cost from O(L²d) to O(L′²d), where L′ ≪ L. After attention, the compact output is written back into the positions it came from inside a zero-initialized tensor shaped like the original sequence. Unselected positions stay zero, which behaves like a hard mask for that one attention operation, and the residual connection carries forward the information of the tokens that were skipped.

The crucial trick is that this masking is per-layer and per-head, not permanent. Since the full sequence dimension is restored every layer, a token ignored in layer 5 can be picked again in layer 6 or in a different head of layer 5. The authors also found that sparsifying every layer hurts, so they first measure how much each layer's hidden states change relative to the layer before (representation drift) and apply sparsification only to the most stable half of layers (δ = 0.5), determined once per model.

Evaluation uses RULER, InfiniteBench, LongBench, and Needle-in-a-Haystack on LLaMA-3.1-8B-Instruct and Mistral-Nemo-12B-Instruct, with τ = 0.005 for LLaMA and τ = 0.008 for Mistral, on a single NVIDIA A100 80GB GPU. All methods are applied only during prefill; decoding uses standard dense attention.

Why This Matters

Impact on research: The paper reframes the sparse-attention design space. Rather than treating block-level sparsity and token-level sparsity as competing approaches, it shows they operate at different granularities and compose. It also makes a methodological point that matters beyond this specific method: irreversible early decisions about token importance are poorly matched to the layer- and head-wise dynamics of real models, and reversibility is cheap to add because the residual stream already preserves the full sequence. The drift metric offers a reusable, model-specific diagnostic for deciding where sparsification is safe.

Real-world applications:

  • Long document summarization, where a single prompt may exceed 100K tokens and the attention cost dominates serving latency.
  • Multi-turn reasoning and agentic workflows that accumulate long conversation histories.
  • Code generation and repository-level code understanding, where context comprises many files.
  • Retrieval-oriented workloads — the RULER and InfiniteBench task sets include retrieval-style questions, multi-hop tracing, key-value retrieval, dialogue, and math, all of which the method exercises.

Industry relevance: The method plugs into existing FlashAttention and block-sparse kernels without modification, which lowers the cost of adoption for teams with production inference stacks. Experiments run on a single A100 80GB, and 128K context is now a standard serving window. Since prefill latency is a dominant cost for long prompts and the method's overhead is under 11% of attention latency, the accuracy-latency trade-off it improves maps directly onto serving cost.

Future Directions

  1. Extending reversibility to decoding. All experiments apply sparsification only during prefill, with decoding left as standard dense attention. Whether interleaved token selection helps or interferes with the memory-bandwidth-bound decode phase — where KV cache eviction methods operate — is unexplored here.

  2. Better-performing or cheaper token importance estimation. The ablation shows scoring choice materially changes accuracy (84.95% to 87.02%), and the current design relies on recent queries. Learning the scoring function, or refining it over layers, is a natural next step.

  3. Learned or finer-grained layer selection. Sparse layers are chosen by a one-time preprocessing step using a drift threshold of δ = 0.5. The appendix's analysis of Low/Mid/High drift groups appears to be cut off in the provided text, so how sensitive accuracy is to δ remains an open question the paper does not fully answer here.

  4. Composition with quantization and other compression axes. The related work surveys quantization and KV cache reduction, but the experiments here compose only with block-sparse attention. Whether token-level sparsification stacks additively with quantization or channel pruning is untested.

Target Audience

Researchers and engineers working on efficient LLM inference, long-context serving, or attention kernel design. It is most useful to readers already comfortable with attention internals who want a practical, kernel-compatible acceleration that composes with what they already deploy, and to those studying where token importance lives across layers and heads. Readers looking for a first introduction to sparse attention will find the motivation section (token-importance dynamics) accessible, while the algorithmic details assume intermediate background.

Authors’ abstract

The quadratic complexity of attention remains the central bottleneck in long-context inference for large language models. Prior acceleration methods either sparsify the attention map with structured patterns or permanently evict tokens at specific layers, which can retain irrelevant tokens or rely on irreversible early decisions despite the layer-/head-wise dynamics of token importance. In this paper, we propose Token Sparse Attention, a lightweight and dynamic token-level sparsification mechanism that compresses per-head $Q$, $K$, $V$ to a reduced token set during attention and then decompresses the output back to the original sequence, enabling token information to be reconsidered in subsequent layers. Furthermore, Token Sparse Attention exposes a new design point at the intersection of token selection and sparse attention. Our approach is fully compatible with dense attention implementations, including Flash Attention, and can be seamlessly composed with existing sparse attention kernels. Experimental results show that Token Sparse Attention consistently improves accuracy-latency trade-off, achieving up to $\times$3.23 attention speedup at 128K context with less than 1% accuracy degradation. These results demonstrate that dynamic and interleaved token-level sparsification is a complementary and effective strategy for scalable long-context inference.

Read the original paper