A lot of prior work addressed key-value (KV) cache selection and compression by sparse attention to enable long-context inference for transformer language models without excessive hardware budgets. We provide a new method for fine-tuning models with sparse attention. It works for any KV cache policy, runs on a moderate hardware budget (e.g., a single Nvidia A100 GPU with 40 GB RAM), and allows the model to co-adapt with the policy, often outperforming models trained with exact attention (sequence parallelism). We also provide an efficient implementation of H2O sparse attention (the leading policy in our experiments) with dedicated scaled dot product attention kernel support. KeysAndValues (https://github.com/awslabs/keys_values), a new open source library for long-context inference and fine-tuning, provides easy-to-use and performant code for all methods discussed here.
Figures & tables
Figure 1: (a) Structure of computation graph as lattice of chunks. The first (prefill) chunk in a row has full length NC , subsequent chunks have length S≪NC . Chunks are grouped into cells (see Section A.5.1 ). (b) autograd is run separately on each cell (l,c) . This needs checkpointed inputs from bottom and left, and head gradients from top (checkpointed) and right (on GPU). The direction is right to left, and head gradients for layer inputs are written back to CPU. See Section A.5.2 for details and notation. (c) Delta encoding of KV cache buffers during autograd on a cell. During the forward pass, deltas of the buffer sequences are extracted, matched in pack_hook calls, and written to the computation graph. During the backward pass, the buffer sequences are reconstructed from the deltas and supplied to unpack_hook . See Section A.5.3 for details.
64k datasets
128k datasets
nq
tri_qa
hot_qa
pop_qa
nq
tri_qa
hot_qa
pop_qa
us
sp
us
sp
us
sp
us
sp
us
sp
us
sp
us
sp
us
sp
exact
-
48.8
-
78.7
-
53.3
-
59.3
-
52.5
-
55.2
-
47.7
-
45.7
lr2k
32.7
14.7
56.8
48.5
43.3
5.7
44.8
28.5
21.2
9.0
55.2
23.5
27.7
5.0
31.0
37.3
slr2k
32.3
15.7
58.2
43.7
43.3
6.3
41.7
24.7
19.5
9.0
50.7
20.2
23.3
7.3
30.8
37.3
h2o2k
41.0
18.2
68.2
54.2
35.7
8.7
47.3
41.2
16.3
7.2
39.5
26.5
17.0
6.0
35.8
35.0
Table 1: Results for long-context inference on Helmet datasets. In the upper part, we use 4 datasets of sequence lengths 64k, 128k each (cols), run 5 KV cache policies and 2 chunk sizes 2048=2k,1024=1k (rows). The MatchFirstWordOrPhrase metric is used with these QA datasets. In the lower part, we use 8 additional datasets and run 3 KV cache policies (chunk size 1024). Different metrics are used, depending on the dataset (see Table 2 ). The exact rows are for exact inference (sequence parallelism). Columns us are for models trained using our novel method with the same cache policy in place, columns sp are for models trained with sequence parallelism.
Appendix figures & tables9 assets
Supplementary material from the paper’s appendix.
Appendix
Category
ID
Source
Metric
Dev
Eval
RAG
nq
Natural Questions
MatchFirstWordOrPhrase
893
600
trivia_qa
TriviaQA
MatchFirstWordOrPhrase
876
600
pop_qa
PopQA
MatchFirstWordOrPhrase
192
600
hotpot_qa
HotpotQA
MatchFirstWordOrPhrase
787
300
Many-shot ICL
trec_coarse
TREC
Accuracy
1000
500
trec_fine
TREC
Accuracy
1000
500
Appendix
Table 2: Overview of the 10 Helmet tasks. Dev and Eval denote the number of instances in the training and evaluation partitions, respectively, at a single context-length setting.
min_len
max_len
train
val
test
min_len
max_len
train
val
test
10150
16160
20
1
4
99524
118799
20
1
4
16231
19351
20
1
4
118916
129020
20
1
4
19548
24601
20
1
4
130291
149101
20
1
4
24603
28956
20
1
4
149721
170360
20
1
4
29279
36905
20
1
4
174598
210951
20
1
4
37408
45690
20
1
4
212022
254268
20
1
4
Appendix
Table 3: Buckets used for stratified random split of LongBench V2 dataset, spaced by 5% percentiles w.r.t. sequence length. A bucket contains sequences of token length in [min_len,max_len] . There are 20 buckets, 19 of size 25, one of size 28. Sequences in a bucket are randomly partitioned into train , val (validation), and test . There are 402 training, 20 validation, and 81 test cases.
64k datasets
128k datasets
nq
tri_qa
hot_qa
pop_qa
nq
tri_qa
hot_qa
pop_qa
lr2k
0.0
1.2
1.3
0.0
0.2
2.3
2.3
0.2
slr2k
1.3
16.2
7.0
5.0
0.5
10.2
3.3
0.8
h2o2k
17.3
58.8
19.7
25.7
4.5
30.5
6.7
9.8
h2o2kno
15.5
48.8
17.0
23.7
3.0
6.0
4.7
2.3
h2o2kor
17.8
56.7
20.3
26.2
7.0
40.0
9.3
9.7
Appendix
Table 4: Results for long-context inference on Helmet datasets. Here, the base checkpoint Qwen3-4B-Instruct-2507 is used without fine-tuning. In the upper part, we use 4 datasets of sequence lengths 64k, 128k each (cols), run 5 KV cache policies and 2 chunk sizes 2048=2k,1024=1k (rows). In the lower part, we use 8 additional datasets and run 3 KV cache policies (chunk size 1024=1k ).
64k datasets
nq
tri_qa
hot_qa
pop_qa
slr128
37.3
61.8
41.7
42.2
h2o128
41.3
72.2
45.3
47.7
h2o128no
41.7
64.8
44.7
46.8
h2o128or
41.5
71.3
38.3
49.2
qh2o2k
30.2
65.8
35.3
42.3
Appendix
Table 5: Results for long-context inference with setups not covered in the main text. We show MatchFirstWordOrPhrase values on test splits for different Helmet datasets nq, hotpot_qa , limiting sequence lengths to 64k tokens. slr128 , h2o128 , h2o128no , h2o128or use chunk size S=128 . qh2o2k and qh2o2kno are variants of Q-Hitter [ 81 ] .
64k datasets
nq
tri_qa
hot_qa
pop_qa
lr2k
113.89 (09.07)
116.89 (07.81)
117.20 (11.45)
106.92 (12.64)
slr2k
116.44 (09.19)
119.54 (07.94)
118.77 (11.50)
106.94 (12.16)
h2o2k
117.93 (09.71)
121.20 (08.09)
118.65 (11.73)
107.93 (12.80)
h2o2kno
116.19 (09.40)
118.93 (08.04)
118.79 (11.60)
107.30 (12.43)
h2o2kor
116.49 (09.52)
118.99 (08.11)
119.74 (11.90)
106.95 (12.44)
Appendix
Table 6: Running time figures for training update step, for Helmet 64k datasets (columns), 5 KV cache policies and chunk sizes 2048=2k,1024=1k,128 (rows). Batch size 8, running on 4 devices. The step from 2k to 1k is 8% to 10% more expensive for lr,slr , 9% to 11% more expensive for h2o variants. The step from 2k to 128 is 136% to 147% more expensive for lr,slr , 151% to 165% more expensive for h2o variants. The step from slr to h2o is 1% to 3% more expensive.
128k datasets
nq
tri_qa
hot_qa
pop_qa
exact
258.38 (15.14)
266.05 (11.31)
262.53 (8.18)
236.74 (21.44)
lr2k
324.02 (22.06)
331.44 (16.58)
331.07 (22.63)
309.37 (24.89)
slr2k
327.93 (21.14)
333.46 (16.24)
335.46 (22.67)
310.74 (24.54)
h2o2k
330.83 (21.98)
338.41 (16.72)
336.79 (23.13)
315.27 (25.83)
h2o2kno
331.48 (22.10)
343.78 (16.75)
337.77 (23.12)
317.60 (25.16)
Appendix
Table 7: Running time figures for training update step, for Helmet 128k datasets (columns), 5 KV cache policies and chunk sizes 2048=2k,1024=1k (rows). Batch size 8, running on 4 devices. The step from 2k to 1k is 11% to 14% more expensive for lr,slr , 13% to 15% more expensive for h2o variants. The step from slr to h2o is 1% to 4% more expensive.
trn
slr1k
h2o1kno
h2o1kor
R
p128
R
p128
R
p128
nq
us
1.1 ± 0.8
0.0 ± 0.0
1.1 ± 1.0
0.0 ± 0.0
1.1 ± 2.7
0.2 ± 4.1
sp
35.5 ± 24.1
99.5 ± 7.1
35.5 ± 21.6
100.0 ± 0.0
36.2 ± 22.1
99.3 ± 8.1
no
35.2 ± 22.0
97.7 ± 15.1
35.8 ± 22.8
99.2 ± 9.1
34.9 ± 22.0
97.8 ± 14.6
trivia_qa
us
1.2 ± 0.8
0.0 ± 0.0
1.3 ± 1.0
0.0 ± 0.0
1.1 ± 0.7
0.0 ± 0.0
sp
35.4 ± 24.1
97.3 ± 16.1
38.2 ± 25.2
99.2 ± 9.1
42.0 ± 24.4
96.7 ± 18.0
Appendix
Table 8: Token length statistics of generated samples for 10 Helmet datasets (of context width 128k) and 3 cache logics. trn denotes model checkpoint being used: us uses our novel method with the same cache policy in place, sp is using sequence parallelism, no is the base checkpoint Qwen3-4B-Instruct-2507 (no fine-tuning). R is based on the ratio of output length to target length (in tokens), p128 (in percent) is the fraction of outputs of maximal size 128 (means, and stddevs over all test set samples).
nq
tri_qa
hot_qa
pop_qa
trec_c
R
p128
R
p128
R
p128
R
p128
R
p128
1.1 ± 1.1
0.0 ± 0.0
1.0 ± 0.6
0.0 ± 0.0
1.1 ± 0.6
0.0 ± 0.0
1.0 ± 0.4
0.0 ± 0.0
1.0 ± 0.0
0.0 ± 0.0
nlu
clc150
inf_qa
inf_mc
json_kv
R
p128
R
p128
R
p128
R
p128
R
p128
1.0 ± 0.1
0.0 ± 0.0
1.0 ± 0.0
0.0 ± 0.0
1.2 ± 0.9
0.0 ± 0.0
1.0 ± 0.0
0.0 ± 0.0
1.0 ± 0.0
0.0 ± 0.0
Appendix
Table 9: Token length statistics of generated samples for 10 Helmet datasets (of context width 128k) for training and inference with exact attention (sequence parallelism).
64k datasets
nq
tri_qa
hot_qa
pop_qa
us
sp
us
sp
us
sp
us
sp
exact
-
55.5
-
85.0
-
64.0
-
61.0
lr2k
40.3
44.8
68.3
74.7
51.3
64.0
46.5
55.8
slr2k
38.0
45.2
72.2
74.5
52.3
65.3
44.2
52.2
h2o2k
46.7
61.3
72.2
85.8
51.0
69.0
52.3
45.8
Appendix
Table 10: Results obtained under the same conditions as Table 1 (upper left), except that the SubEM metric is used instead of MatchFirstWordOrPhrase .
Long contexts have become standard in pretrained LLMs, yet they remain expensive to run: prefill compute grows quadratically with sequence length, and every decode step re-reads a key-value cache that grows linearly with it. Sparse attention cuts these costs by attending only to a relevant subset of past tokens, but selecting that subset is itself expensive. We present SpotAttention, a lightweight selector that attaches to a frozen pretrained transformer and learns by KL distillation to estimate its attention distribution. The selector picks the top-K keys each query attends to, and because its estimate is a calibrated distribution, a dual top-p rule reads the per-query, per-layer budget directly from it. Across Qwen3 (dense, 4B-32B) and Qwen3.5 (hybrid linear/full attention, 4B-9B), SpotAttention matches dense accuracy at contexts up to 128K tokens, eight times the training length. Decode at L=128K runs 3.9x faster than FlashAttention and 1.8x faster than Twilight, the strongest training-free baseline. Quantizing the selector's K-cache to INT4 or FP4 microscale shrinks it 3.5x at no accuracy cost.
Long-context inference in large language models is bottlenecked by the quadratic cost of full attention. Existing efficient alternatives often rely either on native sparse training or on heuristic token eviction, creating an undesirable trade-off among efficiency, training cost, and accuracy. In this work, we show that full-attention LLMs are already intrinsically sparse and can be transformed into highly sparse models with only minimal adaptation. Our approach is built on three observations: (1) only a small subset of attention heads truly requires full long-context processing; (2) long-range retrieval is governed primarily by a low-dimensional subspace, allowing relevant tokens to be retrieved efficiently with a 16-dimensional indexer; and (3) the useful token budget is strongly query-dependent, making dynamic top-p selection more suitable than fixed top-k sparsification. Based on these insights, we propose RTPurbo, which retains the full KV cache only for retrieval heads and introduces a lightweight token indexer for sparse attention. By exploiting the model's intrinsic sparsity, RTPurbo achieves sparsification with only a few hundred training steps. Experiments on long-context benchmarks and reasoning tasks show that RTPurbo preserves near-lossless accuracy while delivering substantial efficiency gains, including up to a 9.36× prefill speedup at 1M context and about a 2.01× decode speedup. These results suggest that strong sparse inference can be obtained from standard full-attention training without expensive native sparse pretraining.
Long-context Large Language Model inference is severely bottlenecked by the massive Key-Value (KV) cache, yet existing sparse attention methods often suffer from static fixed-budget (Top-k) retrieval or rely on proxy scores that are computationally expensive and biased. To address these limitations, we propose RaBitQCache, a novel sparse attention framework that utilizes randomized rotated binary quantization and high-throughput binary-INT4 arithmetic to efficiently estimate attention weights. Our proxy score serves as an unbiased estimator with a proven error bound, enabling adaptive Top-p retrieval that dynamically adjusts the token budget based on actual attention sparsity. We further implement a hardware-aware system with asynchronous pipelining and lazy updates to mask overhead. Evaluations demonstrate that RaBitQCache significantly accelerates inference and reduces memory I/O while preserving generation quality compared to state-of-the-art baselines. Code is available at https://github.com/Sakuraaa0/RaBitQCache.git.
Wenhao Li, Jinhao Dong, Hailin Zhang +3
School of Information, Renmin University of China, Beijing, China · Peking University, Beijing, China