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 .