Learning how to Forget: Fine-tuning for Long-Context Sparse Attention

arXiv:2608.19920 · cs.CL · Submitted 2026-08-20 · Read on arXiv

Listen

Radio episode about this paper

Transcript

Introduction to the show: ident: AI Radio. Generated commentary on the latest Artificial Intelligence papers.

Tom: I'm Tom, and with me are Jane, Lu, senior AI researcher at Tsinghua, Meng, lead engineer at a mysterious AI startup and Lalam, the in-house Large Language Model.

Jane: Today's paper: "Learning how to Forget".

Tom: A novel method for fine-tuning transformer language models with sparse attention is introduced, demonstrating that this approach can be effective on moderate hardware budgets and often outperforms exact attention methods.

Jane: First, who's behind it and why it matters.

Paper summary: Tom: So, to wrap up what we've discussed about "Learning how to Forget: Fine-tuning for Long-Context Sparse Attention," the paper presents a novel training method that handles sparse attention while adapting to any KV cache policy, operating on moderate hardware budgets. It really focuses on improving the fine-tuning process itself.

Jane: And I think the authors are pointing toward a future where we don't have to choose between fast inference and efficient training when working with very long contexts. They show that this new method often outperforms models trained using exact attention techniques like sequence parallelism across several benchmarks.

Lu: The implications for the research community are big because they provide a unified framework for fine-tuning sparse attention, regardless of the specific cache selection strategy you might want to use later on. This flexibility is quite valuable.

Meng: From an implementation standpoint, the contribution from KeysAndValues, which provides performant code and support for things like KV cache buffer quantization and CPU offloading, makes this immediately useful for engineers looking to deploy these models efficiently in real environments.

Lalam: For me, the cultural impact is that this work suggests a path toward building more robust AI systems that can handle extremely long, complex inputs seamlessly. It moves us closer to an AI that doesn't just process short snippets but truly understands the whole narrative context efficiently.

Tom: Absolutely, Lalam. The ability to co-adapt with the cache policy is something we need as models get bigger and documents get longer and longer; this paper gives us a tool for that adaptation during training. That's what it's all about, moving beyond just inference optimizations to improving the core training loop.

Conclusion: Tom: So, we're talking about a paper that tackles how these large language models manage their memory when dealing with really long sequences. Jane, can you explain what the authors actually did in simple terms?

Jane: Well, essentially they developed a training method for sparse attention that works no matter which cache policy you choose later on. They combine some clever tricks to process long sequences using constant resources, which is pretty neat for handling big inputs.

Lu: I find the idea of co-adapting with the KV cache policy during training really interesting; it opens up so many creative avenues for how models can manage their internal state efficiently as context grows.

Meng: From my side, I'm wondering about the practical side—how much memory does this actually save on real hardware compared to what we're running now? We need to know if this is just theoretical or something that actually runs smoothly in production.

Lalam: For me, the big picture here is how this allows AI systems to truly grasp and retain long, complex information without losing focus on the important parts; it feels like a step toward a more coherent form of knowledge processing.

Tom: Exactly! It's about making sure that when we train these massive models, they aren't just guessing about how to store their context; they learn the best way to manage it dynamically. Jane, what are your thoughts on the authors and why this specific focus on fine-tuning sparse attention is so important right now?

Jane: I think the authors really zeroed in on a common bottleneck: getting good performance from sparse attention when you have long contexts. They show that their technique doesn't need extra approximations beyond what we already expect from sparse methods.

Lu: Their focus on the "nested activation checkpointing" and "autograd saved tensors packing" is a very elegant way to manage the gradient computation, which I think shows a deep understanding of how modern hardware constraints interact with model training loops.

Meng: That sounds technically sound, but I'm curious about the trade-offs they mention regarding hardware budgets; if it requires special kernel support or complex setup, does that make it only accessible to big labs?

Lalam: The implication is that we might see AI systems that handle massive amounts of text or data with a much more sustainable and less resource-intensive training process moving forward.

Tom: That's the core excitement here! It suggests we can push the limits of context length without needing exponentially more hardware just to fine-tune effectively. Jane, if you had to summarize the main takeaway for our listeners in one simple sentence, what would you say?

Jane: I'd say this paper gives us a reliable way to train models that are good at handling long contexts using sparse attention, regardless of the specific cache strategy we implement later.

Tom: Pretty powerful statement there! So, we’ve seen how they handle the mechanics and the results; now we need to think about what this means for how AI actually evolves in practice. We're going to look next at what these results imply for real-world applications and potential future research directions.

Matthias Seeger, Zeyu Zhang, Vihang Patil, Konstantinos Benidis, Sebastian Schelter

Amazon Web Services

cs.CL

Submitted: 2026-08-20

Updated: 2026-09-28

Code: https://github.com/awslabs/keys_values

Importance score: 81/100

The gist: A novel method for fine-tuning transformer language models with sparse attention is introduced, demonstrating that this approach can be effective on moderate hardware budgets and often outperforms

Key concepts

Sparse Attention
Instead of calculating attention for every pair of tokens (quadratic complexity), sparse attention only calculates attention for a select subset of relevant token pairs. This drastically reduces the computational cost and memory needed, making it feasible to process very long sequences.
KV Cache Policy
This refers to how the model manages and stores previously computed key and value vectors (the KV cache) during sequence generation. The new method is designed to work seamlessly with any of these policies, allowing the training process to adapt effectively without needing complex re-implementations.
Activation Checkpointing
This is a memory-saving technique where only some parts of a neural network's computation graph are stored during the forward pass. Instead of saving all intermediate results, it saves checkpoints, allowing the model to recompute necessary parts during the backward pass, significantly reducing GPU memory usage.
Sequence Parallelism
This is an alternative training strategy where a single long sequence is split across multiple GPUs. The paper compares its new sparse attention fine-tuning method against this technique, showing that their proposed method often achieves better accuracy while avoiding the long output artifacts seen in sequence parallelism.

Terminology

Summary

A novel method for fine-tuning transformer language models with sparse attention is introduced, demonstrating that this approach can be effective on moderate hardware budgets and often outperforms exact attention methods.

The gist: A new method works for any KV cache policy and requires no further approximations beyond sparse attention. As we demonstrate in experiments on a range of long-context benchmarks, our training algorithm allows the model to co-adapt with the KV cache policy, often outperforming models trained with exact attention (sequence parallelism).

New Fine-Tuning Method

The authors propose a new method for fine-tuning models with sparse attention that works for any KV cache policy and runs on resources comparable to sparse attention inference. This method combines nested activation checkpointing and CPU offloading with exploiting a linear KV cache buffer recurrence by way of autograd saved tensors packing. The goal is to process sequences of arbitrary length with constant resources.

Improvements to Heavy-Hitter Oracle (H2O)

The paper details methodological and implementation improvements for the heavy-hitter oracle (H2O) sparse attention policy. Key improvements include:

  1. Providing Triton code to return summed attention weights alongside a FlashInfer SDPA kernel.

  2. Introducing the normalized H2O score: ϕt h2o-norm(b, h, j) = (t − t(b, h, j))−1ϕt h2o(b, h, j). This score is used to determine where new content is written by overwriting slots with the smallest values.

Gradient Computation Strategy for Sparse Attention

Training models with sparse attention requires addressing GPU memory constraints when computing gradients for sequences of length N much larger than the cache length NC. The method involves several steps:

  1. Avoiding differentiation through the KV cache policy by storing all KV cache policy decisions in a replay log.

  2. Using activation checkpointing [25] in a nested fashion, partitioning chunks into cells to reduce GPU memory requirements from O(L·S−1(N−NC)·NC·D) to O(NC·D) per autograd call.

  3. Exploiting the linear recurrence between neighboring cache buffers by storing delta key, delta value instead of the full keys and values in the computation graph using autograd saved tensors hooks (packing) and reconstructing them during backward pass (unpacking).

Experimental Validation

Experiments on a range of long-context benchmarks show that the training algorithm often outperforms models trained with sequence parallelism. The comparison is conducted across various KV cache policies, including:

((

(lastrec (lr): Keeps the last recent NC − β and first β tokens in the cache.

(smart lastrec (slr): Variant of lastrec where β is chosen dependent on content.

(h2o (h2o), h2o norm (h2ono), h2o orig (h2oor): Variants of H2O.

The results indicate that the proposed method ("us) often outperforms sequence parallelism (sp) on many datasets, particularly for tasks like trec coarse, nlu, clinc150, inf qa, inf mc, json kv, where the metric is mostly Accuracy. However, a consistent failure mode of sequence parallelism is observed: its outputs are far too long and contain mostly random nonsense."

Implementation and Open Source Library

The work is supported by the open source library KeysAndValues (https://github.com/awslabs/keys values), which provides easy-to-use and performant code for all methods discussed here, including support for quantization of KV cache buffers and integration of optimized attention kernels like FlashInfer SDPA. The library also supports CPU offloading of KV cache buffers and model weights.

Kernel Support Needs

The paper highlights a gap in existing fast SDPA kernels: they do not natively support the required operations for sparse attention, such as returning summed attention weights P i mb,h,i,j. The authors suggest that kernel developers should consider extensions like returning these summed weights or allowing for "implicitly defined causal masks of the form (b, h, i, j) 7→ (−∞)I[P +i<t(b,h,j)]" to better support sparse attention inference.

Future Directions

Future work plans include combining context parallelism with sparse attention and exploring kernel fusion ideas in order to narrow the latency gap further. There is also consideration for multi-stream asynchronous implementations which allow for on-the-fly CPU offloading. The authors hope that the KeysAndValues library will make it easier to explore new KV cache policies and approximations.

Summary of Contributions

The main contributions are:

**"New method for fine-tuning transformer language models with sparse attention and arbitrary KV cache policy in place...

Improvements for AI systems

As a fastidious researcher, I have analyzed this paper, Learning how to Forget: Fine-tuning for Long-Context Sparse Attention, and identified several high-impact areas for improving AI systems. The core contribution is a method that allows fine-tuning large language models (LLMs) with sparse attention policies to run on moderate hardware budgets, often outperforming exact attention methods.

Here are the specific improvements and capabilities this research enables:


)Improved AI System Capabilities:

The proposed method allows for the fine-tuning of LLMs using sparse attention mechanisms (like H2O) on moderate hardware (e.g., a single A100 GPU with 40GB RAM), achieving performance comparable to, or better than, models trained with exact attention methods like sequence parallelism.

Specific improvements and capabilities include:

Fine-tuning Models with Arbitrary KV Cache Policies: The method is universally applicable to any existing Key-Value (KV) cache policy (e.g., H2O variants). This means developers are no longer restricted to a single, pre-defined attention approximation; they can choose the eviction/selection logic that best suits their specific downstream task or hardware constraints.

Co-adaptation of Model and Cache Policy: The training algorithm allows the model to co-adapt with the KV cache policy during fine-tuning, often leading to superior performance compared to models trained with exact attention (sequence parallelism). This capability means the model learns not just how to process information, but also how to strategically utilize its limited short-term memory buffer.

Resource Efficiency for Long Contexts: The technique enables training on resources comparable to sparse attention inference. This drastically lowers the barrier for fine-tuning models designed for very long contexts (e.g., 64k or 128k tokens), making state-of-the-art long-context fine-tuning accessible to researchers and smaller labs with limited GPU budgets.

Efficient Gradient Computation via Nested Activation Checkpointing: It introduces a sophisticated, nested activation checkpointing strategy combined with CPU offloading and delta encoding of KV cache buffers. This allows for the computation of gradients over sequences much longer than the available GPU memory would normally permit, effectively reducing the memory requirement for autograd calls from being proportional to the full KV cache size to being proportional only to a small fraction of it (e.g., O(NC · D) per call).

Support for Efficient Policy Implementation: The paper provides methodological and implementation improvements for leading sparse attention policies, specifically the Heavy-Hitter Oracle (H2O). This includes providing optimized Triton code and integrating dedicated Scaled Dot Product Attention (SDPA) kernels, which significantly reduces the latency penalty associated with using complex scoring mechanisms like H2O during training.

Enabling Flexible Inference for Any Context Width: The system can be used for inference on arbitrary context widths (N), as memory requirements remain independent of N, provided the cache length NC is set appropriately. This contrasts sharply with sequence parallelism, which is strictly limited by the physical GPU memory and number of devices available.

Versatile Implementation via KeysAndValues Library: The development of an open-source library, KeysAndValues, provides a clean abstraction layer for all discussed methods (sparse attention policies, quantization, FlexAttention integration) and model checkpoints (via LitGPT). This allows researchers to easily experiment with new sparse attention policies without having to rewrite low-level CUDA kernels for every policy variant.

Mitigation of Inference Failure Modes: Experimental results show that the proposed method outperforms sequence parallelism on many benchmarks, suggesting it better handles the loss of stopping generation or output nonsense failure modes observed in sequence parallelism, indicating a more robust training regime that aligns better with how sparse attention inference actually operates during generation.

Sources

Related papers