Training Transformers for KV Cache Compressibility
Yoav GelbergYam EitanMichael BronsteinYarin GalHaggai Maron
Introduces KV-Compression Aware Training (KV-CAT), a pretraining method that masks key-value slots during training to produce representations that substantially improve the performance of downstream KV cache compression algorithms on long-context tasks.
Deploying autoregressive transformer language models for long-horizon tasks—such as code repository analysis, document synthesis, and autonomous agent operations—is severely bottlenecked by the Key-Value (KV) cache. During inference, models must store key and value states for every token across all attention heads and layers. Memory requirements and decoding latency scale linearly with sequence length, making long context windows a major operational constraint. While post-hoc context and cache compression methods have emerged to mitigate this burden, they operate on fixed, pretrained models whose internal representations may be inherently resistant to compression.
The article establishes both theoretically and empirically that KV cache compressibility is a learned structural property of the model itself rather than a trait of the input data alone. Its primary objective is to demonstrate that transformers can be guided during training to learn internal representations that are substantially more amenable to downstream compression, thereby improving inference efficiency without degrading standard model accuracy.
To achieve this, the authors propose KV-Compression Aware Training (KV-CAT), a continued pretraining framework. Under KV-CAT, the transformer undergoes dual forward passes: a dense pass and a masked pass where lightweight, learned linear-attention routers enforce a train-time KV sparsification policy by masking a set fraction of KV slots. The training objective combines a self-distillation loss that forces the masked representations to mirror the dense output, an anchoring next-token prediction loss to preserve base capabilities, and a budget constraint loss. The authors evaluated the approach on Qwen 2.5 models (0.5B and 1.5B parameters) across standard benchmarks, long-context question answering, needle-in-a-haystack retrieval, and suffix-continuation perplexity under post-hoc compression techniques.
Empirical findings confirm the effectiveness of this approach across several key operational dimensions. First, KV-CAT models fully preserved baseline capabilities, matching standard short-context multiple-choice benchmark accuracy within 0.5 to 0.7 percentage points without compression applied. Second, when downstream compression methods were applied, KV-CAT checkpoints demonstrated up to a 3.21-fold improvement in suffix perplexity retention and reached equivalent performance levels up to 5 times faster during gradient-based cache optimization. Third, in retrieval tasks from compressed contexts, KV-CAT raised mean retrieval accuracy by 5.2 to 6.4 percentage points overall, with improvements reaching 11 to 19 percentage points at moderate retention budgets (30% to 50% keep ratios). Finally, on long-context question answering tasks from LongBench v2, KV-CAT delivered an average accuracy improvement of up to 39% across multiple domains.
These results demonstrate that the bottleneck of post-hoc cache compression can be addressed at the training stage. For engineering and deployment leaders, this offers a practical pathway to substantially lower the operational compute and memory footprint of long-context applications, potentially reducing serving costs and hardware constraints for high-throughput language model pipelines.
Organizations developing or hosting long-context language models should consider incorporating compression-aware objectives, such as KV-CAT, during domain adaptation or continued pretraining phases. When adopting this method, engineering teams can choose between fixed heuristic sparsification policies and adaptive learned routers; the evidence indicates that learned routers offer the best trade-off by preserving uncompressed baseline accuracy while maximizing downstream compressibility.
Decision-makers should note several operational boundaries and uncertainties. KV-CAT introduces non-trivial training overhead and implementation complexity due to the dual forward passes and auxiliary routing modules. Furthermore, empirical evaluations were conducted on smaller open-weight models (up to 1.5B parameters) using specific continued pretraining corpora. While confidence in the reported experimental improvements is high, broader deployment across larger-scale frontier models and specialized domains will require targeted pilot validation.
- Paper: KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV Cache, Zirui Liu et al. (2024). KIVI provides a concrete post-hoc KV-cache quantization baseline, clarifying the compression methods that KV-CAT aims to make more effective through training.
- Paper: Dynamic Memory Compression: Retrofitting LLMs for Accelerated Inference, Piotr Nawrot et al. (2024). Dynamic Memory Compression establishes continued pretraining as a way to adapt models for KV-cache reduction, a useful precursor to KV-CAT’s training-time approach.
No sufficiently relevant recommendations were found.
