KVQuant: Towards 10 Million Context Length LLM Inference with KV Cache Quantization
Coleman Hooper, Sehoon Kim, Hiva Mohammadzadeh, Michael W. Mahoney, Yakun Sophia Shao, Kurt Keutzer, Amir Gholami
Introduction
Large language models (LLMs) have revolutionized many natural language processing (NLP) tasks. In order to improve the capabilities of LLMs, there is significant interest in increasing the context lengths of LLMs. Longer context lengths enable new applications, including long document summarization, retrieval for answering questions about long documents, extended multi-turn applications (chen2023longlora, ), and code analysis. To support this pull from applications, there have been significant recent advances in long-context length models in industry (anthropic, ; openai, ), as well as in academia (chen2023longlora, ).
Given the importance of LLM workloads, there is strong motivation to improve their inference efficiency. LLM inference with large context lengths can be incredibly resource-intensive; serving LLMs requires high-end GPUs, and the largest LLMs require costly multi-GPU inference setups. When analyzing the computational nature of generative inference with LLMs, it becomes quickly apparent that, for relatively small batch sizes, the computation is memory bound (kim2023squeezellm, ). With the growing divergence between computational speeds and memory speeds, this problem is only going to get worse over time (gholami2020ai_and_memory_wall, ). This makes reducing the memory bottleneck preeminently important. Further analysis shows that the memory bottleneck is strongly related to context size. For short sequence lengths, the dominant contributor to memory consumption is the weight matrices, and therefore the optimal strategy is to minimize the model size in order to reduce memory consumption as well as bandwidth requirements (kim2023full, ; kim2023squeezellm, ). However, for long sequence lengths, the main bottleneck is the memory requirements for caching Key and Value (KV) activations throughout inference. In particular, the size of the KV cache can become the dominant contributor to memory footprint, even for a 32K context limit (see Table 1), making it challenging to perform long context length inference. This challenge is further exacerbated when one considers batched inference.
It is therefore crucial to develop methods for compressing the KV cache to enable efficient long-sequence length inference. Existing approaches lead to unacceptable accuracy degradation due to the outlier structures in KV cache activations as well as suboptimal bit allocation with existing uniform and non-uniform approaches. In this work, we perform an extensive analysis of KV cache activations in recent LLMs, revealing patterns which can be exploited to enable ultra-low precision quantization with minimal accuracy loss. In particular, we make the following contributions (summarized in Figure 1):
We find that the Key matrices exhibit structured outliers in specific channels before applying RoPE. However, the outlier channel magnitudes become less consistent after applying RoPE, posing a distinct challenge for low precision quantization. Based on these observations, we use per-channel quantization for Keys, and we quantize Keys before RoPE is applied (see Section 3.1 and Section 3.2).
We find that existing uniform and non-uniform quantization methods result in sub-optimal quantization signpost placement. Instead, we propose a Non-Uniform Quantization (NUQ) method which considers sensitivity and not just magnitude when quantizing activations. We show that one can apply sensitivity-weighted non-uniform quantization offline on a calibration set to derive accurate datatypes for KV cache quantization (see Section 3.3).
Even with the above, we find that outlier values in cached KV activations can significantly degrade quantization resolution. Unlike for weights, it is non-trivial to extract outlier values from activations, given the dynamic nature of activations. However, we find that we can efficiently and accurately identify and compress outlier values in order to store them compactly in a separate sparse representation. We also find that per-vector outlier detection outperforms per-matrix outlier detection with no additional memory overhead. With this method, we can attain under 0.1 perplexity degradation for 3-bit KV cache quantization on both Wikitext-2 and C4 by only removing 1% of outliers, thereby facilitating accurate inference with 4.8 longer context length.
For ultra low-bit precision, we find that the quantized activations can deviate significantly from their corresponding fp16 values. To address this, we propose a lightweight Q-Norm layer which shifts and scales the distribution after de-quantization to match the mean and standard deviation of corresponding fp16 values. Interestingly, the Q-Norm layer can be fused with non-uniform quantization values resulting in no overhead during inference. This was particularly helpful for 2-bit quantization (see Section 3.5).
We implement custom CUDA kernels to perform activation quantization efficiently during inference, achieving up to 1.4 speedups for Key and Value matrix-vector multiplications for LLaMA-7B at 4-bit precision relative to the fp16 baseline (see Section 3.7 and 4.2). These results demonstrate how our methodology allows for accurate and efficient low-bit KV cache quantization.
Background
When inferring a decoder-only LLM, inference proceeds in two distinct phases. In the prefill phase, the model takes in an input prompt, which it processes in parallel. During the generation phase, the model then generates the output sequence autoregressively, meaning that each token generation is dependent on all previously generated tokens. As such, for small batch sizes, the generation phase of LLM inference is typically memory-bandwidth bound, as the only available parallelism is across different sequences in a given batch.
Additionally, during generation, the model needs to store intermediate Key and Value activations in order to condition generations on previously generated output tokens. Otherwise, we would need to recompute all prior Keys and Values at each timestep, which would be prohibitively expensive. For each prior token, we need to store the Keys and Values at each layer in order to use these activations when generating future tokens. These stored activations are referred to as the Key-Value (KV) cache. Throughout this paper, we will capitalize Key and Value to distinguish when we are referring to the KV cache tensors. Assuming a model with layers and attention heads with dimension , the KV cache size for batch size and sequence length is , meaning that it grows linearly with both batch size and sequence length. As shown in Table 1, the KV cache becomes the dominant contributor to memory consumption for longer sequence lengths and larger batch sizes. Note that since each sequence in batched inference depends on separate past context, there is no available batch-level parallelism when loading the cached Keys and Values for their respective computations in batched inference. KV cache loading is therefore always memory-bandwidth bound. This motivates pursuing methods to optimally compress the KV cache, even at the expense of a more complex dequantization process.
2. LLM Quantization
KV Cache Quantization. There have been many prior works on LLM quantization. Several have focused on weight-only quantization for LLMs, due to the greater contribution to memory consumption and runtime for fairly small sequence length and batch size (lin2023awq, ; dettmers2023spqr, ; kim2023squeezellm, ). There has also been work on quantizing both weights and activations (including KV cache) (xiao2023smoothquant, ; shao2023omniquant, ). However, there is still a significant perplexity degradation when quantizing KV cache activations to low precision; (sheng2023flexgen, ; zhao2023atom, ) quantized KV cache activations to 4-bits, but required fine-grained grouping for 4-bit quantization, while still observing some perplexity degradation, and (sheng2023flexgen, ) observed that 3-bit KV cache quantization with fine-grained grouping leads to unacceptable accuracy loss. Other works quantized KV cache activations to 4-bits but required retraining to maintain performance (liu2023llmqat, ). One concurrent work also explores low precision KV cache quantization in order to enable larger batch size inference by reducing the KV cache size (kivi, ). In this work, we introduce a method for near-lossless low-bit KV cache quantization that minimizes performance degradation without the need for finetuning.
Outlier-Aware LLM Quantization. LLMs have been known to have distinct outliers both in weights and activations (dettmers2022llm, ; dettmers2023spqr, ; kim2023squeezellm, ). SqueezeLLM and SpQR both decompose the weight matrix into a sparse matrix containing a small portion of outliers and a dense matrix that can be accurately quantized to low precision (referred to as dense-and-sparse or sparse-quantized representation) (dettmers2023spqr, ; kim2023squeezellm, ). LLM.int8() (dettmers2022llm, ) handled particular outlier channels separately in higher precision, and SmoothQuant (xiao2023smoothquant, ) migrates quantization difficulty due to outlier channels to weights in order to support joint weight-activation quantization. Other works reconsidered the dimension along which we quantize in order to reduce quantization error (or else added per-channel compensation to improve quantization performance) (bondarenko2021understanding, ; heo2023rethinking, ; wei2022outlier, ; wei2023outlier, ). In this work, we demonstrate that per-channel pre-RoPE Key quantization provides significant accuracy benefits given the outlier structure in Keys, and that dense-and-sparse quantization can be efficiently applied for KV cache quantization.
Non-uniform LLM Quantization. Non-uniform quantization has also been applied in the context of LLMs. Non-uniform quantization allows for more flexible quantization signpost placement relative to uniform quantization methods, enabling improved accuracy for the same bit precision (kim2023squeezellm, ; dettmers2023qlora, ). Building on the observation that model parameters tend to be approximately normally-distributed, prior work has proposed the NormalFloat datatype (dettmers2023qlora, ). SqueezeLLM (kim2023squeezellm, ) derived per-channel non-uniform quantization signposts using a sensitivity-weighted k-means approach. In this work, we show that we can derive accurate per-layer non-uniform datatypes using a sensitivity-weighted k-means approach with KV cache activations.
3. KV Cache Compression
There have also been several prior works on compressing the KV cache. Some of these methods aim to only store important tokens in the KV cache and to evict less important tokens, thereby maintaining low memory usage (liu2023scissorhands, ; zhang2023h, ; ge2023model, ). Other methods aim to only retrieve a subset of tokens at each step to achieve memory bandwidth savings (ribar2023sparq, ). In this work, we explore KV cache quantization as an orthogonal direction for compressing the KV cache in order to enable long context inference.
Method
To inform our approach, we first performed a detailed analysis to understand the KV cache distributions. Figure 2 shows sample distributions for the KV cache activations. We observe that the Key matrices tend to have distinct outlier channels, which have larger average magnitudes than other channels; this corroborates previous observations about outlier channels in LLM activations (dettmers2022llm, ; xiao2023smoothquant, ). The Value matrices exhibit both outlier channels as well as outlier tokens (although these outliers are less extreme than the outlier Key channels).
Existing KV cache quantization approaches perform per-token quantization (meaning that the scaling factor and zero-point are shared by elements in the same token) (sheng2023flexgen, ; zhao2023atom, ). However, due to the differing average magnitudes between channels, the values within a channel are easier to quantize when grouped together than the values across different channels. As such, to better match the distributions, we investigate per-channel KV cache quantization, meaning that the scaling factor and zero-point are shared by elements in the same channel. By sharing the scaling factor and zero-point along the channel dimension, this will naturally group together values with similar magnitudes, thereby mitigating the impacts of outlier channels on other channels when quantizing to low precision. As outlined in Appendix D, we find that per-channel quantization provides significant accuracy benefits for Keys but not for Values. By leveraging per-channel quantization for Keys and per-token quantization for Values, we observe a 3.88 perplexity improvement on Wikitext-2 for 3-bit LLaMA-7B quantization. Note that this can potentially add runtime overhead since the quantization dimension is now misaligned with the reduction dimension for the Keys when performing matrix-vector multiplications. However, we find that we are able to efficiently dequantize Keys and perform the Query-Key matrix-vector multiplication without adding runtime overhead, as shown in Section 4.2. Additionally, as outlined in Section 3.6, per-channel quantization can also be challenging due to the need to recompute scaling factors as tokens are added to the Key cache. We show that we can calibrate offline for scaling factors, thereby avoiding expensive online recomputation.
Per-channel Key quantization was also explored in another concurrent work (kivi, ), which leveraged similar intuition about grouping together large magnitude values in the same channel to minimize quantization error. Their methodology requires fine-grained grouping for per-channel quantization while maintaining a residual subset of the KV cache in fp16 (until all elements for that group have been added to the KV cache). In our work, we demonstrate that by leveraging offline calibration, we can accurately perform per-channel quantization without grouping.
2. Pre-RoPE Key Quantization
3. nuqX: An X-Bit Per-Layer Sensitivity-Weighted Non-Uniform Datatype
Uniform quantization is suboptimal for KV cache quantization since the Query and Key activations are non-uniform. Additionally, KV cache loading is memory bandwidth bound, regardless of batch size or sequence length, meaning that the dequantization overhead introduced by non-uniform quantization methods is not problematic (since the added computation does not introduce any additional latency). It is therefore desirable to leverage non-uniform quantization methods for KV cache quantization.
In (kim2023squeezellm, ), the authors computed non-uniform quantization signposts using a sensitivity-weighted k-means approach. However, this is challenging to apply online during inference due to its computational cost, and it is also difficult to estimate sensitivity for activations online. We therefore facilitate efficient online non-uniform KV cache quantization by computing sensitivity-weighted quantization signposts offline on a calibration set prior to inference. Using the diagonal Fisher information matrix (derived in Appendix C), along with the quantization error for activation , we formulate the error minimization objective as:
In order to derive a per-matrix non-uniform datatype, we first normalize each vector to the range $$. We then minimize the objective in Equation 2 offline on a calibration set using a k-means solver in order to obtain the quantization signposts for the non-uniform datatype for each Key or Value layer. Appendix F compares our non-uniform quantization approach with existing uniform and non-uniform quantization baselines (dettmers2023qlora, ), demonstrating how our non-uniform approach provides 0.33 perplexity improvement on Wikitext-2 for LLaMA-7B relative to 3-bit uniform methods. Table 15 in Appendix I shows how computing the required Fisher information for the LLaMA-65B model takes only a few minutes, and how using the k-means solver takes only a few minutes per layer (with the computation for each layer being parallelizable). In Appendix N, we also demonstrate that we can derive a metric for accurate one-shot mixed-precision assignment (where different layers are assigned different bit widths) using the quantization error weighted by sensitivity information.
4. Per-Vector Dense-and-Sparse Quantization
Figure 5 in Appendix B shows the portion of elements falling within different percentiles of the dynamic range. For both Keys and Values, the majority of elements are contained within a small percentage of the dynamic range. This means that by leveraging dense-and-sparse quantization, as demonstrated in (kim2023squeezellm, ), in order to isolate a small percentage of numerical outliers, we can restrict the range that we need to represent, thereby allowing us to represent the remaining elements with greater precision.
Additionally, when looking at the Key and Value distributions in Figure 2, different channels and tokens have different average magnitudes. Therefore, an element which counts as an outlier in one channel may not be an outlier in another channel (since that channel may have a greater average magnitude). It is therefore crucial to directly target the outlier values that skew the dynamic range at the granularity that we are quantizing in order to address the values that are exaggerating the range along that particular dimension. In this work, we leverage per-vector dense-and-sparse quantization, where we use a different outlier threshold per-vector (either a separate threshold per-channel for per-channel quantization, or a separate threshold per-token for per-token quantization), rather than a single outlier threshold for each layer.
Note that computing outlier thresholds for per-vector dense-and-sparse quantization poses potential accuracy and efficiency challenges. However, in Section 3.6, we show that we are able to accurately calibrate for per-channel outlier thresholds offline and efficiently compute per-token outlier thresholds online. After determining the upper and lower outlier thresholds, the remaining numbers in the vector are normalized to the range $$, and we then minimize Equation 2 (ignoring outliers) in order to obtain the quantization signposts for the non-uniform datatype for the remaining numbers. Appendix G will demonstrate the benefits of removing a small percentage of numerical outliers and keeping them in full precision, as well as the advantages of per-vector dense-and-sparse quantization over using a single outlier threshold for each layer. As shown in Figure 1, by removing 1% of numerical outliers using per-vector outlier thresholds, we achieve an additional 0.25 perplexity improvement on Wikitext-2 for 3-bit LLaMA-7B quantization, which is within 0.08 perplexity of the fp16 baseline.
5. Mitigating Distribution Shift using Q-Norm
When pushing to extremely low bit widths like 2-bit quantization, we begin to observe accuracy degradation due to distribution shift, meaning that the post-quantization distribution has a different mean and standard deviation than the pre-quantization distribution. As shown in Appendix H, this distribution shift can lead to greater error accumulation at later layers. To combat this, we introduce Q-Normalization (shortened to Q-Norm), where we normalize the quantization centroids obtained from k-means in order to ensure the post-quantization distribution has the same mean and standard deviation as the pre-quantization distribution. Our solution is inspired by (li2023norm, ), which demonstrated that ensuring that the activation distributions post-weight quantization have similar mean and standard deviation to the activation distributions pre-weight quantization can help reduce accuracy degradation. Given mean and standard deviation for the pre-quantization distribution and mean and standard deviation for the post-quantization distribution (computed offline on a calibration set), we normalize the quantization centroids to as follows:
During dequantization, we use in place of such that the mean and standard deviation of dequantized values matches the original fp16 distribution. As shown in Appendix H, Q-Norm provides significant accuracy benefits for 2-bit quantization with no added inference cost.
6. Offline Calibration versus Online Computation
A crucial challenge for activation quantization is that we either need to compute statistics on-the-fly (which is potentially expensive) or else we need to use offline calibration data (which potentially has negative accuracy implications). The challenges with computing scaling factors (and zero-points) online versus offline for both Keys and Values are shown in Figure 3. In per-channel quantization, it is challenging to update scaling factors online since the scaling factors corresponding to each incoming channel would potentially need to be updated whenever a new token is added to the KV cache. It is therefore desirable to be able to compute statistics offline (i.e., using calibration data before running inference). While this can have negative effects on model accuracy, in Appendix I we show that we can effectively calibrate offline for per-channel quantization, obviating the need for online updates of scaling factors for per-channel quantization. For per-token quantization, it is challenging to calibrate for scaling factors offline due to the presence of outlier Value tokens. It is therefore desirable to be able to compute scaling factors and outlier thresholds online for each incoming token. As shown in Appendix I, we can efficiently compute outlier thresholds online per-token by offloading to the CPU. By leveraging custom quantization function implementations for compressing activations, we are able to perform online per-token Value quantization without compromising on performance.
7. Kernel Implementation
In order to efficiently perform activation quantization on-the-fly, we leverage dedicated kernel implementations with our 4-bit quantization method for compressing vectors to reduced precision and extracting the sparse outliers, performing matrix-vector multiplications using the compressed vectors, and performing sparse matrix-dense vector multiplications using the sparse outliers. We store the quantized Key and Value matrices as 4-bit elements which are used as indices into lookup tables to recover the corresponding fp16 values. We store the sparse outlier matrices in either Compressed-Sparse Row (CSR) or Compressed-Sparse Column (CSC) format (depending on which aligns better with appending new Key and Value tokens). The kernels for the Key matrix-vector operations apply RoPE on-the-fly in order to support pre-RoPE quantization. More kernel implementation details are provided in Appendix O.
Results
We used the LLaMA-7B/13B/30B/65B, LLaMA-2-7B/13B/70B, and Mistral-7B models to evaluate our methodology by measuring perplexity on both Wikitext-2 and C4 (touvron2023llama, ; touvron2023llama2, ; jiang2023mistral, ); see Appendix J for details on our experimental setup. Table 2 shows the results for LLaMA models for the Wikitext-2 dataset. We compared our method with per-token quantization with and without grouping. The baseline configurations used by Atom and FlexGen are included for reference (zhao2023atom, ; sheng2023flexgen, ). We did not include results for KIVI (kivi, ) since they did not report perplexity values. We find that our method consistently outperforms baseline approaches by an especially large margin with 3-bit and 2-bit quantization. Once we incorporate outliers, we further push the performance of low-precision quantization, achieving 4-bit quantization with less than 0.02 perplexity degradation, 3-bit quantization with under 0.09 perplexity degradation, and 2-bit quantization with under 0.43 perplexity degradation relative to the fp16 baseline across all models and datasets (while attaining 3.7, 4.8, and 6.9 memory savings, respectively). Tables 16 and 17 in Appendix K show full perplexity evaluation on Wikitext-2 and C4, showing the consistent performance of our approach across different models and datasets. Table 18 in Appendix K also provides zero-shot MMLU evaluation with 3-bit quantization, demonstrating how our method maintains performance on downstream tasks.
1.2. Long Context Length Evaluation
We evaluated long context length performance using the LLaMA-2-7B-32K model (uptrained for long sequence lengths using positional interpolation (chen2023extending, )) as well as the LLaMA-2-70B-32K LongLoRA model (chen2023longlora, ). For evaluating performance on longer context lengths, we first evaluated perplexity on Wikitext-2 using larger amounts of input context, as shown in Figure 4 (chen2023longlora, ; han2023lminfinite, ). The results demonstrate how our method maintains accuracy even for longer amounts of input context, thereby enabling efficient and accurate long sequence length inference. Additional long sequence length perplexity evaluation results for 2-bit quantization (including Q-Norm) are provided in Appendix M.
We also evaluated the performance of our quantization method on passkey retrieval to assess the model’s ability to use its context. Passkey retrieval involves evaluating the model’s capacity to locate specific information in long texts (longchat2023, ), and this can be used to effectively measure the maximum distance over which a token can attend during the inference stage. We used the passkey evaluation framework from (zhu2023pose, ) (which is based on the methodology from (mohtashami2023landmark, )) to evaluate retrieval performance. Table 3 shows passkey retrieval results for the LLaMA-2-7B-32K and LLaMA-2-70B-32K models. These results indicate that with 4-bit and 3-bit quantization, KVQuant is able to maintain retrieval performance of the fp16 model. We observe some degradation with 2-bit quantization especially with LLaMA-2-70B, indicating potential room for improvement in 2-bit quantization techniques.
1.3. Joint Weight and KV Cache Quantization
Table 4 shows results for our KV cache quantization method when the weights are also quantized using the methodology in SqueezeLLM (kim2023squeezellm, ). We observe minimal perplexity degradation when leveraging our KV cache quantization approach, even when weights are also quantized to reduced precision (achieving within 0.02 perplexity of 4-bit weight-only quantization when quantizing the KV cache using nuq4-1% for the LLaMA-7B and LLaMA-13B models). These results demonstrate how our method is compatible with existing weight-only quantization methods.
2. Performance Analysis
Table 5 shows kernel benchmarking results using a batch size of 1 for the 4-bit dense-and-sparse compression and matrix-vector kernel implementations. We show results across different sequence lengths to assess the performance of the kernels at different points during generation. We report latency benchmarked on an A6000 GPU and averaged over 1000 runs. The results show that for the Key and Value multiplications, we can achieve 1.1-1.2 and 1.2-1.4 latency savings, respectively, relative to the baseline. We have integrated these kernels into an end-to-end generation pipeline that is able to compress activations dynamically during inference, thereby achieving significant memory savings and allowing for either larger batch sizes or longer sequence lengths.
3. Pushing the Context Length Limit
Table 6 shows the KV cache memory requirements for 128K, 1M, and 10M sequence lengths, with the KV cache stored in fp16 as well as 4-bit, 3-bit, and 2-bit precision with KVQuant. As one can see, our method provides 3.7 KV cache compression (nuq4-1%) and enables serving the quantized LLaMA-65B model with a context length of 32K tokens on a single A100-80GB GPU (requiring 30GB for the model weights compressed to 4-bit, and 33GB for the KV cache when compressed with nuq4-1%), and even serving the LLaMA-7B model with a context length of 1M tokens on a single A100 GPU (requiring 3GB for the model weights in 4-bit precision and 66GB for the KV cache with nuq4-1%). Additionally, when considering an 8-GPU serving system, we enable serving the LLaMA-7B model with 10M context length (with nuq3-1%) or the LLaMA-65B model, with 1M context length (with nuq4-1%). While currently there are no models or datasets to measure accuracy with such large context sizes, our results on smaller context lengths show little degradation compared to baseline inference, demonstrating the benefits of our approach for enabling long sequence length inference.
Conclusion
As context lengths in LLMs increase, the KV cache activations surface as the dominant contributor to memory consumption. Quantization is a promising approach to reduce the size of KV cache activations, but prior solutions failed to represent activations accurately in ultra-low precisions, such as sub-4-bit. In contrast, we achieve accurate ultra-low precision KV cache quantization. By quantizing Keys per-channel before applying RoPE, we are able to better match the outlier distribution and mitigate the impacts of RoPE on quantization (due to it mixing pairs of channels which may have different average magnitudes). We use non-uniform quantization to better allocate the small number of quantization signposts at low precision. We observe significant accuracy improvements when employing dense-and-sparse quantization, particularly when detecting outliers at the same granularity as we compute quantization scaling factors. Crucially, we demonstrate that we can perform accurate calibration offline for Keys, as well as efficient online scaling factor and outlier threshold computation for Values. By leveraging these methods, we are able to enable accurate low-precision activation quantization, achieving 4.8x compression (nuq3-1% outliers) with only 0.1 perplexity degradation across different LLaMA, LLaMA-2, and Mistral models. Additionally, we show that Q-Norm improves accuracy for 2-bit quantization by helping mitigate distribution shift. Our methodology therefore supports inferring the LLaMA-7B model with a context length of 10M on an 8-GPU serving system. Through our efficient kernel implementation, we are able to show improved latency relative to the fp16 baseline, demonstrating how our method allows for improved latency in addition to the memory savings.
Limitations
While our work enables accurate long-context length inference by reducing the memory requirements, there is significant work required for training long context length models with greater than 100K context length. This work is orthogonal to our efforts, which are constrained to efficient inference with long context length models. Additionally, our latency benchmarking results currently focus on memory-bandwidth bound generation rather than prompt processing (where we need to compress multiple Keys and Values at once). In future work, we plan to develop dedicated efficient kernels for block Key/Value compression in order to efficiently compress multiple Keys and Values at once. Finally, in the current end-to-end implementation, there are inefficiencies in how memory allocation is handled for updating the sparse matrix (where the data corresponding to the previous tokens have to be copied when concatenating them with the data from the new token). In future work, we plan to optimize this by doing blocked allocation to avoid overheads from reallocating memory.
Acknowledgements.
The authors would like to acknowledge Nicholas Lee for helpful discussions and feedback. We acknowledge gracious support from Intel, Furiosa, Apple, and Samsung SAIT. We also appreciate the support from Microsoft through their Accelerating Foundation Model Research, including great support from Sean Kuno. Furthermore, we appreciate support from Google Cloud, the Google TRC team, and specifically Jonathan Caton, and Prof. David Patterson. Prof. Keutzer’s lab is sponsored by the Intel corporation, Intel One-API, Intel VLAB team, the Intel One-API center of excellence, as well as funding through BDD and BAIR. We appreciate great feedback and support from Ellick Chan, Saurabh Tangri, Andres Rodriguez, and Kittur Ganesh. Sehoon Kim would like to acknowledge the support from the Korea Foundation for Advanced Studies (KFAS). Amir Gholami was supported through funding from Samsung SAIT. Michael W. Mahoney would also like to acknowledge a J. P. Morgan Chase Faculty Research Award as well as the DOE, NSF, and ONR. Our conclusions do not necessarily reflect the position or the policy of our sponsors, and no official endorsement should be inferred.
References
Appendix A RoPE Equation
By leveraging this element-wise implementation, we can apply RoPE on-the-fly to the Key activations (after dequantizing the Key activations and before multiplying them with the corresponding elements in the Query vector).
Appendix B Key and Value Dynamic Range
Figure 5 shows the portion of the elements contained within difference percentages of the dynamic range for both Keys and Values. The majority of values ( 99%) are contained in a small portion of the dynamic range, and a small portion of numerical outliers skew the dynamic range that must be represented. This motivates our dense-and-sparse approach which removes numerical outliers and stores them in a separate sparse matrix, thereby restricting the range that needs to be represented in the dense component.
Appendix C Derivation for Sensitivity Analysis
The following derivation is adapted from , with the only difference being that it is estimating sensitivity with respect to activations rather than gradients. Let the activation at a given layer for input data be denoted as , and let the loss of the neural network with respect to activation be represented as . Let the gradient and Hessian of the loss function with respect to activation be denoted as and , respectively. Let be the application of the quantization function to , and let represent the quantization error. Note that the subsequent derivations assume that the activation and gradient matrices are all flattened to one dimension.
By performing Taylor expansion of the loss with respect to the Hessian for data , we get:
Assuming that the loss has converged to a local minimum, can be approximated as zero, which (omitting constants) yields the following formula for the increased error in the loss due to perturbation in :
Appendix D Per-Channel Key Quantization Ablations
We report results in Table 7 demonstrating the perplexity for different KV cache compression ratios, showing that per-channel quantization for Keys and per-token quantization for Values outperforms the standard per-token quantization approach for both Keys and Values, yielding an improvement of 3.88 perplexity for the LLaMA-7B model at 3-bit precision. This demonstrates the benefits of per-channel Key quantization to mitigate the large outlier channels in Keys. Additionally, although there are per-channel outliers in Values, we observe that per-channel quantization for Values actually performs worse than per-token quantization. We hypothesize that this behavior is because per-channel Value quantization leads to greater error accumulation in particular output values (since the result of the attention scores multiplied by one channel of the Values will be localized to a single value in the output vector), which leads to greater quantization error at later model layers. Another concurrent work, KIVI , observes similar behavior for per-channel Value quantization, which they attribute to the fact that per-token Value quantization confines the error to each token. Assuming that the output is a weighted sum of only a few important tokens (as only a few attention scores are large), a perturbation in these tokens can lead to significant degradation.
Appendix E Pre-RoPE Key Quantization Ablations
As shown in Table 8, pre-RoPE Key quantization achieves higher accuracy than post-RoPE quantization, with an improvement of 0.65 perplexity for 3-bit quantization with the LLaMA-7B model. These results show that the rotary positional embeddings make Key quantization more challenging due to mixing pairs of channels with different magnitudes. Pre-RoPE quantization thereby allows for more accurate quantization at low precision.
Appendix F Sensitivity-Weighted Non-Uniform Quantization Ablations
Table 9 shows perplexity evaluation results across different LLaMA, LLaMA-2, and Mistral models on Wikitext-2 using nf3, nuq3, and nuq3 without using sensitivity-weighting. These results demonstrate the benefits of our sensitivity-weighted non-uniform quantization approach, relative to NormalFloat quantization , as we achieve consistent accuracy improvements of up to 0.32 perplexity across different models (with particularly pronounced improvements for larger models). For 3-bit quantization with the LLaMA-7B model, we observe a 0.33 perplexity improvement relative to uniform quantization. The gains relative to uniform quantization are particularly noticeable for 3-bit and 2-bit quantization, where the benefits of non-uniform quantization are more pronounced due to the reduced precision. These results also demonstrate the necessity of our sensitivity-weighting approach in order to derive performant non-uniform datatypes using a k-means based approach.
Figure 6 shows the relationship between the Fisher information and the normalized activation values for the LLaMA-7B model. For the Value matrices, we observe that the average Fisher information is higher close to the middle, which leads to quantization centroids being pulled closer to the center of the distribution. For the Key matrices, we observe that the average sensitivity across different magnitudes is relatively constant, which leads to wider spacing for the centroids. These results show that using Fisher information to derive a sensitivity-weighted per-layer datatype allows for better representation of Keys and Values.
Appendix G Per-Vector Dense-and-Sparse Quantization Ablations
Table 10 shows the performance improvements we observe when isolating a small portion of outliers and storing them in a sparse format. We provide results both with using a single per-matrix outlier threshold, as well as with applying separate outlier thresholds per-vector. In particular, we see greater improvements by employing outlier detection with a different threshold per-channel for Keys and per-token for Values. This provides additional benefits since some values which would be considered outliers for the entire matrix are not actually outliers within a particular channel (so they are not hard to quantize). It is therefore better to directly target the outliers that will skew the quantization range for a particular channel. By removing 1% of outliers using per-vector thresholds, we can achieve an additional 0.25 reduction in perplexity for the LLaMA-7b model at 3 bits, thereby enabling 3-bit quantization with under 0.1 degradation in perplexity.
Appendix H Q-Norm Ablations
Table 11 provides perplexity results for LLaMA-7B and LLaMA-13B with and without Q-Norm. As shown in the table, Q-Norm provides noticeable accuracy improvements for 2-bit quantization, but it doesn’t provide significant accuracy improvements for 3-bit or 4-bit quantization. Figure 7 demonstrates how the minimization of distribution shift provided by Q-Norm leads to reduced quantization error (particularly at later layers in the network).
Additionally, we experimented with using per-vector Q-Norm, where we normalize each channel for Keys and each token for Values to ensure the distribution has the same mean and standard deviation post-quantization. For Keys, per-vector Q-Norm requires offline computation of the required normalization per-channel; and for Values, per-vector Q-Norm requires online computation of the required normalization per-token. With per-vector Q-Norm, the normalization parameters must also be stored separately from the per-vector outlier thresholds and cannot be fused with the NUQ datatype. Table 11 also compares employing per-vector Q-Norm with per-matrix Q-Norm. Although we observe similar accuracy with per-vector Q-Norm when incorporating dense-and-sparse quantization, we observe worsened performance when we aren’t using dense-and-sparse quantization. We attribute this to the larger relative shift when applying normalization per-vector, as well as the impacts that significant changes to the centroids can have on outlier values (in the case where we aren’t using dense-and-sparse quantization, large perturbations in these outlier values can lead to significant accuracy loss). The need to only slightly tweak normalization parameters when combating distribution shift is similar to , which describes how weight quantization can be improved by slightly adjusting layernorm parameters to avoid distribution shift. Due to the improved performance of per-matrix Q-Norm for dense-only quantization (as well as potential inference overheads with per-vector Q-Norm from having to separately rescale the non-uniform centroids per-vector and compute normalization statistics on-the-fly for Values), we focus on per-matrix Q-Norm. However, there is potential to improve accuracy through fine-grained normalization in future work; for example, Figure 8 demonstrates how per-vector Q-Norm provides improved performance for the LLaMA-2-70B-32K model when evaluating perplexity on long context lengths.
Appendix I Calibration Ablations
Table 12 shows accuracy results when using offline calibration for computing the scaling factors for the Keys. For 4-bit quantization, we observe no accuracy loss when calibrating scaling factors offline. For 3-bit quantization, we observe minor accuracy degradation when not employing outlier extraction methods. However, if we remove a small percentage of outliers, then the accuracy with offline calibration is the same as computing the scaling factors online per-channel during evaluation. This demonstrates that when incorporating outlier extraction methods, we are better able to perform offline calibration due to reduced sensitivity to outliers (either to outliers during calibration that exaggerate the quantization range, or to outliers during evaluation that cannot be represented accurately if there weren’t large outliers observed during calibration).
Table 13 shows the runtime for the operation for the LLaMA-7B model (which is required for computing outlier thresholds online). It compares the runtime of the operation with the runtime for the QKV projections, finding that the runtime is 60% of the matrix-vector operation runtime for a single projection layer. The operation can also be performed efficiently on the CPU, so we can actually run this operation in parallel with the subsequent linear layer matrix-vector operations on the GPU (which is possible by computing the Value projection before the Key and Query projections). This allows us to compress the activations dynamically without added runtime overhead, thereby enabling online scaling factor computation for the Value tensors.
Table 14 shows the perplexity of the LLaMA-7B model using different numbers of samples during calibration. The results show that perplexity is similar across the range of the number of samples tested for each bit width. This shows that the calibration step does not require a large number of calibration samples to attain high accuracy. Additionally, Table 15 shows how both Fisher information computation and calibration (including k-means) per-layer take only a few minutes for the LLaMA-65B model on a typical server machine. Even if we perform calibration sequentially for each layer, the entire calibration process would take a maximum of 6 hours for the LLaMA-65B model at 4-bit precision.
Appendix J Additional Experimental Details
For our empirical evaluation, we use 16 calibration samples of sequence length 2K from the Wikitext-2 training set (as well as the corresponding gradients) to derive the per-channel scaling factors and zero-points, to derive the non-uniform datatypes for both Keys and Values, and to estimate layer-wise sensitivity for mixed-precision experiments. We measured perplexity on both Wikitext-2 and on C4 using a sequence length equal to the maximum context length of the model (2K for LLaMA, 4K for LLaMA-2, and 8K for Mistral-7B). For baseline experiments, we use post-RoPE quantization, both since this is required from an efficiency perspective without a dedicated kernel implementation, and because it provides better accuracy when quantizing Keys per-token as shown in Appendix L.
We make several assumptions in order to estimate average bit widths and KV cache sizes for different approaches. We compute these estimates assuming a sequence length of 128K. For integer quantization, we assume a low-precision integer offset and a 16-bit scaling factor, whereas for NormalFloat and NUQ we assume that the zero-point and offset are each 16-bit. For the sparse matrices, 32-bit integers are assumed for the per-token indices (since we need to support long sequence lengths), and the elements and per-element indices are assumed to be 16-bit. This means that for CSR, the rows are assumed to be 32-bit and the columns and values are assumed to be 16-bit, whereas for CSC, the columns are assumed to be 32-bit and the rows and values are assumed to be 16-bit.
Appendix K Full Perplexity Evaluation and MMLU Evaluation
Tables 16 and 17 show perplexity evaluation results across different LLaMA, LLaMA-2, and Mistral models on Wikitext-2 and C4, respectively. These results demonstrate the benefits of our approach for KV cache compression across different model sizes as well as across different language modeling datasets. Table 18 also provides zero-shot evaluation results on MMLU for LLaMA-7B and LLaMA-13B, demonstrating how we can maintain accuracy even for 3-bit KV cache quantization . We used the Language Model Evaluation Harness to run zero-shot evaluation across all MMLU tasks .
Appendix L Post-RoPE Per-Token Quantization Ablation
Table 19 shows perplexity evaluation on Wikitext-2 for the LLaMA-7B model with uniform quantization, with Keys quantized pre-RoPE and post-RoPE. These results show that post-RoPE Key quantization is superior to pre-RoPE Key quantization when quantizing Keys per-token. This is because when rotating an outlier channel with large average magnitude and another channel with smaller average magnitude together, at some positions in the sequence part of the magnitude from the outlier channel will be shifted to the smaller channel. This partially mitigates the impact of the outlier channel on skewing the quantization range for some of the tokens in the sequence. As such, for our baseline comparisons, we use post-RoPE per-token Key quantization to serve as a stronger baseline.
Appendix M Additional Long Sequence Length Evaluation
Figure 8 shows perplexity evaluation results with varying amounts of input context for 2-bit quantization. These results show that Q-Norm (and in particular, per-vector Q-Norm) can provide significant perplexity advantages for 2-bit quantization with long sequence length models. Additionally, for the LLaMA-2-70B-32K model, we observe perplexity degradation when going to a context length of 32K, likely due to error accumulation; however, this is largely mitigated by employing per-vector Q-Norm.
Appendix N Mixed-Precision Quantization
An additional method to optimize compression performance is to consider mixed-precision quantization, where different layers are assigned different bit widths. This can allow for more accurate compression down to low bit widths due to differing sensitivities to quantization error of different layers. Finding the best bit precision distribution for different layers is intractable with brute force methods as the search space is exponentially large. We therefore aim to derive a metric to determine which layers can be quantized to reduced precision with minimal degradation in model accuracy, thereby enabling efficient one-shot mixed-precision bit assignments.
Prior work on mixed-precision quantization has leveraged the largest Hessian eigenvalue or the Hessian trace to assess which layers are most sensitive to quantization ; however, due to the computational challenges of computing the full Hessian or even estimating Hessian eigenvalues, we instead leverage the Fisher information approximation for the Hessian (as derived in Appendix C). Using the diagonal Fisher information matrix along with the quantization error, we can use the following sensitivity metric for a given layer (with activation and quantized activation ) to encompass both the quantization error as well as the Fisher information for that layer:
We use the quantization error computed at the lower precision (as well as sensitivity information computed in fp16) to determine which layers were most sensitive to being quantized to the lower precision level. Figure 9 shows the mixed-precision perplexity results on the Wikitext-2 dataset for the LLaMA-7B and LLaMA-13B models. We show mixed-precision results using our sensitivity metric, and we include both quantization error-based mixed-precision assignment and the inverse of the selection order from our sensitivity metric as baselines. We see improved performance from incorporating sensitivity-based analysis for determining activation precisions when employing mixed-precision quantization. Our sensitivity metric therefore allows for efficiently determining an accurate mixed-precision assignment in order to trade off KV cache size and model accuracy.
Appendix O Kernel Implementation Details
We implemented 4-bit lookup table-based kernels for matrix-vector multiplication between the Key or Value activations (packed as a lookup table (LUT) plus indices into the LUT per-element) and a full-precision activation vector. These kernels load the compressed Key and Value activations and dequantize them only as needed in order to minimize memory bandwidth utilization. All arithmetic is performed in fp16. The lookup table entries are the values of the sensitivity-weighted non-uniform datatype for that particular layer scaled according to the range of activations that need to be represented . Note that it is possible to implement this using a single LUT shared across channels that is rescaled by a per-channel or per-token scaling factor and offset, but for simplicity we used a separate LUT per-channel or per-token for our initial implementation.
When selecting between the Compressed-Sparse Column (CSC format) and the Compressed-Sparse Row (CSR) format for storing the outliers for the Keys and Values, we needed to consider how easy it would be to append new vectors. When using CSC format for the Key matrix, we only need to append a single element to the column vector, as well as one new element to the row and value vectors per nonzero element in that new column. If we used CSR format, we would need to insert the new column and value elements in the middle of the existing column and value vectors, and we would need to recompute the elements of the row vector. When using CSR format for the Value matrix, we only need to append a single element to the row vector, as well as one new element to the column and value vectors per nonzero element in that new row. If we used CSC format, we would need to insert the new row and value elements in the middle of the existing row and value vectors, and we would need to recompute the elements of the column vector. We therefore used the CSC format for the Key matrices and the CSR format for the Value matrices.
One challenge with efficiently processing the sparse matrix-dense vector operation is that the sparsity distribution may be unbalanced. This poses a challenge for efficiently processing the sparse matrix on a GPU as there can be different numbers of nonzeros to process per thread. We therefore leverage a balanced sparse matrix-dense vector kernel based on , which assigns an equal number of nonzeros per thread. This has greater synchronization overhead than assigning a single thread for an entire row or column when processing CSR/CSC matrices, but it leads to a more balanced work assignment between threads. We set the number of threads such that there were 10 nonzero values assigned to each thread. The dense non-uniform kernel and balanced sparse kernels are launched in one call to avoid overhead from summing the output vectors from these separate operations.
Table 20 shows a detailed breakdown of kernel runtime, including how much time is spent packing vectors into the compressed format and how much time is spent on the dense and sparse matrix-vector multiplications. We find that even with 1% sparsity, we can attain significant speedups of up to 1.4 relative to the fp16 matrix-vector multiply kernels, demonstrating how our methodology facilitates efficient inference with a low-precision quantized KV cache.