A Systematic Study of Cross-Layer KV Sharing for Efficient LLM Inference
You Wu, Haoyi Wu, Kewei Tu
Introduction
A major bottleneck for the deployment of LLMs is memory consumption, of which the key-value (KV) cache in the transformer architecture occupies a large portion Kwon et al. (2023). Various methods have been proposed to reduce the memory consumption of the KV cache in LLMs. For example, Shazeer (2019); Ainslie et al. (2023) share the KVs across query heads and Zhang et al. (2023); Xiao et al. (2024) keep the KV cache of only a small portion of tokens.
More recently, several methods are proposed in which the KVs are computed only at a subset of transformer layers and shared to the other layers, such as LCKV Wu and Tu (2024), YOCO Sun et al. (2024) and CLA Brandon et al. (2024). These methods not only significantly reduce memory consumption but also improve inference speed, while preserving the performance of LLMs in language modeling and downstream tasks. However, while all these methods are based on the idea of cross-layer KV sharing, they differ significantly in how the sharing is done.
In this study, we consider a unified framework for cross-layer KV sharing, of which LCKV, CLA, and YOCO can be seen as special configurations. We then empirically test all the configurations of the framework, including several novel ones that have never been considered in previous work. Our experiments show that all the configurations can achieve significantly higher throughput than the standard transformer when the prompt is short, but the throughput of the configurations that compute the KVs at the top layers degrades dramatically when the prompt is long. On the other hand, while the performance of most configurations is comparable with that of the standard transformer when only half of the layers rely on the KVs computed by the other layers, it is the configurations that compute the KVs at bottom layers whose performance degrades the most when more layers become reliant on the other layers for the KVs. We hope our framework and empirical studies would help users interested in cross-layer KV sharing to choose methods and configurations according to their throughput and performance requirements. Our code is available at https://github.com/whyNLP/LCKV.
Existing Methods
Layer-Condensed KV Cache (LCKV) Wu and Tu (2024) computes the KVs of only the top layer of the transformer, which are paired with queries of all the layers. Consequently, LCKV omits the KV computation and discards the KV parameters for all the layers other than the top layer. To prevent severe performance degradation, LCKV also optionally retains standard attention for a small number of top and bottom layers.
You Only Cache Once (YOCO) Sun et al. (2024) computes the KVs of only the top layer of the transformer, which are paired with the queries of the top-half of the layers. The bottom-half of the layers uses efficient attention to achieve a constant cache size. Goldstein et al. (2024) uses a similar sharing pattern to YOCO, but further compresses the size of the KV cache.
Cross-Layer Attention (CLA) Brandon et al. (2024) uniformly divides transformer layers into multiple groups of adjacent layers. In each group, it pairs the queries of all the layers with the KVs of the bottom layer. Zuhri et al. (2024) shares the KVs in the same way as CLA, but applies a more efficient training scheme. Liu et al. (2024) groups every two adjacent layers in the middle-to-deep portion and compresses the KV cache across each group. Chen et al. (2024) groups non-adjacent layers and pairs the queries of the upper layer with the KVs of the lower layer in each group. Rajput et al. (2024) uses a combination of the sliding window attention layers and a similar sharing pattern to CLA. Liao and Vargas (2024); Mu et al. (2024); Rajabzadeh et al. (2024) apply the sharing pattern similar to CLA to the computed attention weights instead of KVs.
A Unified Framework
Unifying previous methods, we propose a framework for cross-layer KV sharing that can be applied to any transformer-based model. Suppose that the transformer has layers. We denote as the index of the layer whose KVs are paired with the queries of the -th layer. If , then layer is called a KV layer, which computes its own KVs that are paired with its queries just as in a standard transformer. Otherwise, layer does not compute its own KVs and instead uses the KV of layer . In this case, we call layer the target layer of layer . Since layer does not need to compute KVs, it does not need weights . Therefore, the number of KV layers determines the number of weight parameters and hence the size of a transformer model. Below we define different configurations of our framework assuming the number of KV layers always set to .
We define a configuration by partitioning transformer layers and positioning target layer(s) differently. We choose the layer partitioning from { pizza, sandwich, lasagna } and choose the target layer positioning from { bottom, top, middle }We also consider positioning at quarter and three-quarter, which is discussed in Appendix D.. The pizza partitioning sets the first layers as KV layers. The sandwich partitioning sets the first layers and the last layers as KV layers. For the remaining consecutive layers in both pizza and sandwich, their target layer is positioned at either the top, the middle, or the bottom of these layers. The lasagna partitioning uniformly divides the layers into groups of consecutive layers. For each group except the first, the target layer of all the layers within the group is positioned at either the top, the middle, or the bottom of these layers. For the first group, however, we always set the bottom layer as the target layer because we empirically find that there is a significant drop in performance if the first layer is not a KV layer.
Note that for the top and middle positioning of the target layer, there exists a cyclic dependency between the target layer and the lower non-KV layers: for each token, its KVs at the target layer is required for attention computation at lower non-KV layers, but are not computed until computation at all the lower layers is finished. So, we follow Wu and Tu (2024) and drop the attention of each token to itself, which is equivalent to masking the diagonal of the attention matrix in each layer.
Table 1 illustrates all the nine configurations that we have defined. We name each configuration with its partitioning and positioning pattern. The sandwich-top, pizza-bottom and lasagna-bottom configurations correspond to LCKV, YOCOThe pizza-bottom configuration differs from YOCO in that it uses the standard attention instead of the efficient attention for the bottom-half of the layers. and CLA respectively. The lasagna-top configuration and all middle configurations are novel and have not been considered in previous work.
For the bottom positioning, the model can be trained in the same way as a standard transformer model. For the top and middle positioning, however, the attention computation of each token at layer depends on KVs of the previous tokens at its target layer , creating sequential dependencies that spoil parallel training. Following Wu and Tu (2024), we perform iterative training to break the sequential dependencies. In each iteration, we pair the queries of each layer with the KVs of its target layer from the previous iteration. For a token sequence of length , parallel training with iterations is equivalent to sequential training. In order to reduce the training cost, we backpropagate the loss only through the last iterations, and use iterations to approximate the KVs of the first iterations.
Note that not all layers need to be trained iteratively. For some configurations, there exist layers without any sequential dependencies at the top and bottom, and we can compute these layers in one pass before and after iterative training, respectively. Therefore, for the pizza and sandwich partitioning, we perform iterative training only on the layers ranging from the first non-KV layer to its target layer, and for the lasagna partitioning, we perform iterative training only on the layers ranging from the first layer of the second group and the target layer of the last group.
2 Inference
The inference of LLMs can be divided into the prefilling and decoding stages. During the prefilling stage, we can conduct early exit Sun et al. (2024) after computing the KVs of the last KV layer. For the top and middle positioning, we perform parallel encoding of the prompt in spite of sequential dependencies by iterative computation with iterations in the same way as in training. The decoding stage is the same as in a standard transformer.
Experiments
We conduct experiments to compare the generation throughput and performance of the standard Llama baseline Touvron et al. (2023) and the nine configurations with different numbers of KV layers. Our implementation is based on HuggingFace Transformers Wolf et al. (2020) with kernel replacement with FlashAttention 2 Dao (2024), fused RMS norm, fused cross-entropy, and fused SwiGLU. Our experiments are conducted on models with 110M and 1.1B parameters, whose configurations are shown in Appendix A. We set and for the top and middle configurations. The sandwich configurations coincide with the pizza configurations when there are only two KV layers and the lasagna-middle configuration coincides with the lasagna-top configuration when the number of KV layers is half of the total number of layers (i.e., 6 and 11 for the 110M and 1.1B models, respectively), therefore omitted in our experiments.
We test the generation throughput of the standard Llama and the nine configurations with 1.1B parameters on an RTX 3090 (24GB) GPU with different sequence lengths. The evaluation follows the settings of FlexGen Sheng et al. (2023).
Figure 1(a) reports the maximum throughput. When the prompt is short (i.e., 5+2043), the prefilling time can be ignored and the generation throughputs of all the nine configurations are almost identical, which are much higher than the baseline throughput and increase as the number of KV layers decreases. When the prompt is long (i.e., 512+1024), the prefilling time becomes significant for the top and middle configurations because of iterative encoding of the prompt. Consequently, their throughputs degrade dramatically, falling below the baseline in some cases. On the other hand, the bottom configurations still achieve significantly higher throughputs than the baseline because no additional computation for prompt is required.
2 Performance on Small Training Set
We train the standard Llama and the nine configurations with 110M and 1.1B parameters from scratchWe also tried model initialization with pre-trained models, the results of which are shown in Appendix C. on the Minipile dataset Kaddour (2023) with 1.7B tokens for one epoch and two epochs, respectively, and evaluate their perplexity. The training details are shown in Appendix A.
Figure 1(b) reports the perplexity. It can be seen that more KV layers lead to better performance in most cases. When the number of KV layers is half of the total number of layers, the performance of most configurations is comparable with that of the baseline. As we reduce the number of KV layers, the performance degrades for almost all the configurations, but the top and middle configurations are less affected compared to the bottom configurations. Two exceptions are the lasagna-top and lasagna-middle configurations, whose performance usually improves with fewer KV layers. This may be due to the fact that the more KV layers there are, the more difficult it is to accurately approximate all the KVs with iterative training.
It can also be seen that the pizza-bottom and lasagna-bottom configurations perform relatively well among all the bottom configurations, and the sandwich-top and sandwich-middle configurations perform relatively well among all the top and middle configurations, respectively. Therefore, we decide to train these four configurations with more data to further investigate their potential in language modeling and downstream tasks.
3 Performance on Large Training Set
We train the standard Llama and the four well-performing configurations with 1.1B parameters from scratch on a 100B subset of the SlimPajama dataset Soboleva et al. (2023) for one epoch and evaluate their perplexity and downstream task accuracy. The training details are shown in Appendix A. We evaluate the perplexity on a 10M subset of the development set of SlimPajama. We also use the LM Eval Harness framework Gao et al. (2023) to test the zero-shot performance on commonsense reasoning tasks including Hellaswag Zellers et al. (2019), OpenBookQA Mihaylov et al. (2018), WinoGrande Sakaguchi et al. (2021), ARC-Easy and ARC-Challenge Clark et al. (2018), BoolQ Clark et al. (2019), PIQA Bisk et al. (2020), and SciQ Welbl et al. (2017).
Figure 1(c) reports the perplexity and average accuracy of downstream tasks. Detailed results of downstream tasks are shown in Appendix B. It can be seen that the sandwich-top configuration
performs better than the two bottom configurations in both perplexity and downstream task accuracy, except for an outlier of the lasagna-bottom configuration with 7 KV layers in downstream task accuracy. The sandwich-middle configuration performs best when the number of KV layers is small.
Conclusion
In this study, we propose a new framework for LLM cross-layer KV sharing that includes previous methods as special cases. We conduct systematic experiments on various configurations of the framework with different KV cache memory budgets and observe their generation throughput and performance in language modeling and downstream tasks. The experimental results show that the pizza-bottom and lasagna-bottom configurations can reduce the size of the KV cache by without too much performance degradation or introducing additional training and prefilling time. However, if one wishes to further reduce the size of the KV cache, cares less about additional training time, and needs to generate sequences much longer than prompts, then the sandwich-middle configuration may be a better choice.
Limitations
In this study, we only conduct experiments on models with 1.1B parameters and training set with 100B tokens. Due to the limited computational resources, we do not explore the performance of larger models with more training data.
References
Appendix A Model and Training Details
Table 2 and 3 show the model configurations and training details for Section 4. The configuration of the 1.1B model follows that of TinyLlama Zhang et al. (2024). We use the MiniPile Kaddour (2023) (licensed under MIT) and SlimPajama Soboleva et al. (2023) (various licenses depending on the data source) as our datasets. Our use of the datasets is consistent with their intended use.
Appendix B Detailed Downstream Task Results
Table 4 reports the accuracy of each downstream task of the models in Section 4.3.
Appendix C Initializing with Pre-trained Models
Instead of training from scratch, we can initialize the standard Llama and the nine configurations with pre-trained models to get better performance. We follow the uptraining scheme of MLKV Zuhri et al. (2024). For each KV layer, we initialize the weights with the averaged weights of all layers whose queries are paired with its KVs. We use the TinyLlama checkpoint trained on 2.5T tokens to initialize the models with 1.1B parameters. The training details are the same as in Section 4.2.
Figure 2 reports the perplexity. It can be seen that all models achieve better performance, compared to training from scratch. The lasagna-bottom configuration performs best when retaining 11 and 7 KV layers, but was surpassed by some top and middle configurations when retaining 3 KV layers. Notice that for the top and middle positioning, we drop the attention of each token to itself and therefore differ from the standard transformer. In future work, we will try to make up for this gap by specially computing the attention of each token to itself, and we hope to get a better performance.
Appendix D More Options for Target Layer Positioning
In addition to positioning the target layer at the top, bottom, and middle, we also consider the quarter and three-quarter, and name the corresponding configurations as middle-1/4 and middle-3/4. We train the new configurations with 1.1B parameters. The training details are the same as in Section 4.2.
Figure 3 reports the perplexity. We omit lasagna configurations because there are not enough layers in each group to distinguish between different target layer positions. It can be seen that the performance of the middle-1/4 and middle-3/4 configurations mainly lies between the top and middle configurations.