Palu: KV-Cache Compression With Low-Rank Projection

发表时间: 2025-04 · arXiv:2407.21118 (ICLR 2025)

原文: https://arxiv.org/abs/2407.21118

Chi-Chih Chang, Wei-Cheng Lin, Chien-Yu Lin, Chong-Yan Chen, Yu-Fang Hui, Pei-Shuo Wang, Ning-Chi Huang, Luis Ceze, Mohamed S. Abdelfattah, Kai-Chiang Wu
National Yang Ming Chiao Tung University, University of Washington, Cornell University

速读

一句话结论 本文提出了一种名为 Palu 的免微调 KV-Cache 压缩框架,通过低秩投影技术压缩 KV 张量的隐藏维度,在保持强零样本精度的同时实现了 50% 的内存压缩率,并为基于 RoPE 的注意力模块带来了最高 1.89 倍的推理提速。

要解决什么问题 大语言模型在长上下文推理时,KV-Cache 的体积会急剧膨胀,导致显存容量和内存带宽受限,进而拖慢解码阶段的推理速度。现有的免微调 KV-Cache 压缩技术主要分为量化(降低数值位宽)和 token 驱逐(丢弃部分缓存)两条路线,但它们都忽略了 KV 张量在隐藏维度上存在的巨大冗余。如果直接在推理时对缓存矩阵进行低秩投影来压缩隐藏维度,又会带来无法接受的矩阵分解计算开销,导致得不偿失。因此,如何在不引入极高运行时开销的前提下,利用隐藏维度的冗余来进一步压缩 KV-Cache,成为了一个亟待解决的卡点。

怎么做的 Palu 的核心思路是将低秩分解的开销从推理阶段转移到离线阶段。它不对运行时的 KV-Cache 直接做分解,而是利用奇异值分解对模型原有的键和值线性投影权重矩阵进行静态分解,即 $\mathbf{W} \approx \mathbf{A}\mathbf{B}$。在推理的前向传播中,输入 token 先通过降维矩阵 $\mathbf{A}$ 投影为低维隐状态 $\mathbf{h} = \mathbf{A}\mathbf{x}$ 并存入缓存,解码时再通过升维矩阵 $\mathbf{B}$ 实时重建出完整的键或值 $\mathbf{y} = \mathbf{B}\mathbf{h}$。为了进一步消除重建开销,Palu 设计了离线矩阵融合机制:将值的升维矩阵直接吸收到注意力输出投影矩阵中;对于不使用旋转位置编码(RoPE)的模型,将键的升维矩阵吸收到查询投影矩阵中;而对于使用 RoPE 的模型,则通过定制的 GPU 算子在片上内存完成高效的在线重建。在分解粒度上,单头独立分解会导致精度严重下降,而所有头联合分解又会带来极高的重建计算量。为此,Palu 提出了分组多头低秩分解,将若干个注意力头的权重拼接后共同分解,在捕获跨头共享信息与控制计算成本之间取得了平衡。此外,不同网络层对压缩的敏感度不同,Palu 采用基于费舍尔信息的自动秩搜索算法,为重要性高的层分配更大的秩。最后,针对低秩分解会在隐状态中引入严重异常值、从而阻碍低比特量化的问题,Palu 引入了沃尔什-哈达玛变换来平滑分布,并通过公式 $\mathbf{W} \approx (\mathbf{A}\mathbf{R})(\mathbf{R}^T\mathbf{B})$ 将哈达玛矩阵 $\mathbf{R}$ 无缝融合到分解后的权重矩阵中,实现了对量化技术的完美兼容,且零运行时额外开销。

效果如何 实验在 Llama-2、Llama-3、Mistral 和 LongChat 等模型上展开,硬件使用单张 RTX 4090 显卡。对比基线包括未压缩的 FP16 原模型,以及代表量化路线的 Atom(按 token 量化)、KVQuant(非均匀量化加稀疏矩阵处理异常值)和 KIVI(细粒度分组量化)。在 WikiText-2 数据集上,Palu 仅靠低秩投影即可在 50% 压缩率下维持极低的困惑度上升;当叠加 2 比特量化时,Palu 的困惑度比同等极限量化基线 KVQuant 低 1.19,同时额外节省了 30% 到 50% 的显存。在 64K 长上下文设置下,结合 50% 低秩压缩与 4 比特量化,Palu 在基于 RoPE 的注意力模块上实现了 2.91 倍的提速,在非 RoPE 模块上提速高达 6.17 倍,端到端生成延迟分别降低至原模型的约三分之一和五分之一,显著优于 KIVI。不过作者也承认了该方法的局限性:在短上下文(如 4K)场景下,该方法几乎没有加速效果;对于 RoPE 模型,当序列长度超过 16K 时,在线重建的计算成本会急剧上升,导致算子层面的加速比开始衰减;此外,在 LongBench 长文本理解任务中,50% 的低秩压缩率难以完全保持精度,需要将压缩率退回到 30% 才能将精度损失控制在 1% 以内。

主要贡献

大型语言模型(LLMs)在推理过程中,将键值状态(KV-Cache)缓存在内存中是加速推理的有效技术。然而,随着上下文长度的增加,KV-Cache的体积会迅速膨胀,给内存容量和带宽带来巨大压力,并导致解码阶段出现内存受限的性能瓶颈。现有的后训练KV-Cache压缩方法(如量化和词元驱逐)主要关注降低数值位宽或保留部分词元,却忽略了KV张量隐藏维度中存在的巨大冗余。

本文的研究目标是通过压缩KV张量的隐藏维度,减少LLM推理时的内存使用量。为此,作者提出了Palu,这是一个利用低秩投影技术的KV-Cache压缩框架。Palu的创新点主要包括:
1. 设计了静态分解线性层为低秩矩阵的架构,通过缓存压缩后的中间状态并在运行时动态重建完整的键和值,有效降低了内存占用并避免了运行时的分解开销。
2. 提出了一种中等粒度的组头低秩分解(Group-head low-rank decomposition, G-LRD)方案,在模型精度和重建效率之间取得了最佳平衡。
3. 引入了一种基于Fisher信息的高效秩搜索算法,能够根据不同矩阵的敏感度自适应地分配秩大小。
4. 实现了低秩感知与量化兼容的增强设计,通过无缝融合Hadamard变换消除了低秩分解引发的异常值,且不增加任何运行时开销。
5. 开发了经过算子融合优化的GPU内核,大幅提升了推理速度。

Palu 用于减少 KV-Cache 的低秩投影方法。线性投影的权重矩阵 $\mathbf{W}$ 被分解为两个低秩矩阵。输入 $\mathbf{X}$ 被向下投影到潜在表示 $\mathbf{H}$,并被缓存。$\mathbf{Y}$ 可以使用向上投影矩阵 $\mathbf{B}$ 从 $\mathbf{H}$ 重建。
Palu 用于减少 KV-Cache 的低秩投影方法。线性投影的权重矩阵 $\mathbf{W}$ 被分解为两个低秩矩阵。输入 $\mathbf{X}$ 被向下投影到潜在表示 $\mathbf{H}$,并被缓存。$\mathbf{Y}$ 可以使用向上投影矩阵 $\mathbf{B}$ 从 $\mathbf{H}$ 重建。

背景知识

多头注意力机制的基础定义。多头注意力(MHA)机制是Transformer架构的核心组件。给定一个新的输入词元 $\mathbf{x} \in \mathbb{R}^d$,一个具有 $n$ 个头的MHA会使用权重矩阵 $\mathbf{W}_i^q$、$\mathbf{W}_i^k$ 和 $\mathbf{W}_i^v$ 将输入分别投影为每个头 $i$ 的查询、键和值,计算公式为 $\mathbf{q}_i = \mathbf{x}\mathbf{W}_i^q, \mathbf{k}_i = \mathbf{x}\mathbf{W}_i^k, \mathbf{v}_i = \mathbf{x}\mathbf{W}_i^v$。其中,$\mathbf{k}_i$ 和 $\mathbf{v}_i$ 代表头 $i$ 在时间步 $t$ 的键和值。接着,可以计算每个头的注意力分数及相应的注意力输出,公式为 $\mathbf{p}_{t,i} = \mathrm{Softmax}\left(\frac{\mathbf{q}_i\mathbf{K}_i^T}{\sqrt{d_h}}\right), \mathbf{a}_i = \mathbf{p}_i\mathbf{V}_i$,其中 $\mathbf{K}_i$ 和 $\mathbf{V}_i$ 表示第 $i$ 个头当前及所有历史键和值的拼接。最终的MHA输出通过拼接所有头的输出并应用输出投影层 $\mathbf{W}_o$ 得到,即 $\mathbf{MHA}(\mathbf{x}) = \sum_{i=1}^h \mathbf{a}_i\mathbf{W}_i^o = \sum_{i=1}^h (\mathbf{p}_i\mathbf{V}_i)\mathbf{W}_i^o$,其中 $\mathbf{W}_i^o \in \mathbb{R}^{d_h \times d}$ 代表每个头 $i$ 的输出投影矩阵的子矩阵。

奇异值分解(SVD)在低秩近似中的应用。SVD是计算给定矩阵低秩近似的常用技术,Palu默认使用该方法进行低秩分解。对于给定的权重矩阵 $\mathbf{W} \in \mathbb{R}^{m \times n}$,SVD将其分解为三个矩阵:$\mathbf{W} = \mathbf{U}\boldsymbol{\Sigma}\mathbf{V}^T$。这里,$\mathbf{U}$ 和 $\mathbf{V}$ 是分别包含左奇异向量和右奇异向量的正交矩阵,$\boldsymbol{\Sigma}$ 是由奇异值组成的对角矩阵。低秩近似过程可以描述为 $\mathbf{W} \approx \mathbf{A}\mathbf{B}$,其中 $\mathbf{A} = \mathbf{U}_r\sqrt{\boldsymbol{\Sigma}_r}$ 且 $\mathbf{B} = \sqrt{\boldsymbol{\Sigma}_r}\mathbf{V}_r^T$。在此公式中,$\mathbf{A} \in \mathbb{R}^{m \times r}$,$\mathbf{B} \in \mathbb{R}^{r \times n}$,$\boldsymbol{\Sigma}_r \in \mathbb{R}^{r \times r}$ 且包含最大的 $r$ 个奇异值,而 $\mathbf{U}_r$ 和 $\mathbf{V}_r^T$ 是从 $\mathbf{U}$ 和 $\mathbf{V}^T$ 截断得到的对应奇异向量。这种截断和矩阵重组使得可以使用两个低秩矩阵 $\mathbf{A}$ 和 $\mathbf{B}$ 来近似权重矩阵 $\mathbf{W}$,从而将存储需求减少了 $\frac{mr + rn}{mn}$。

方法细节

利用SVD重写线性投影层。为了比在运行时直接分解KV-Cache更高效地应用低秩投影,Palu使用SVD对键(Key)和值(Value)的投影矩阵进行静态分解。这一方法基于一个核心观察:低秩分解可以将线性投影层从 $\mathbf{y} = \mathbf{x}\mathbf{W}$ 重写为 $\mathbf{y} = \mathbf{x}\mathbf{A}\mathbf{B}$。在这里,$\mathbf{A} \in \mathbb{R}^{d \times r}$ 是低秩投影矩阵,而 $\mathbf{B} \in \mathbb{R}^{r \times d}$ 是通过SVD推导出的重建矩阵。前向传播过程首先将输入词元 $\mathbf{x} \in \mathbb{R}^d$ 向下投影到一个低维的潜在空间 $\mathbf{h} \in \mathbb{R}^r$ 中,然后再将其向上投影回原始空间,即 $\mathbf{h} = \mathbf{A}\mathbf{x}, \mathbf{y} = \mathbf{B}\mathbf{h}$。

缓存与重建的两步过程。这种分为两步的处理过程使得Palu能够实现两个关键目标:首先,模型可以存储较低维度的潜在表示,而不是原始的键和值状态;其次,模型可以在解码期间动态地将这些潜在表示重建为完整的键和值。

Palu 使用低秩分解 ($\mathbf{W} \approx \mathbf{AB}$) 将键(或值)投影到较低维度的潜在表示 $\mathbf{h}$,从而减小 KV-Cache 的大小。原始键 ($\mathbf{K}_t$) 使用 $\mathbf{B}^k$ 动态重建,而 $\mathbf{B}^v$ 被融合到 $\mathbf{W}^o$ 中以避免重建开销。这种融合也减轻了输出投影的计算负担。
Palu 使用低秩分解 ($\mathbf{W} \approx \mathbf{AB}$) 将键(或值)投影到较低维度的潜在表示 $\mathbf{h}$,从而减小 KV-Cache 的大小。原始键 ($\mathbf{K}_t$) 使用 $\mathbf{B}^k$ 动态重建,而 $\mathbf{B}^v$ 被融合到 $\mathbf{W}^o$ 中以避免重建开销。这种融合也减轻了输出投影的计算负担。

注意力头的键值矩阵分解。在将Palu与注意力机制结合时,系统会对键和值的线性层进行分解。对于每一个注意力头 $i$,Palu应用SVD技术,将键投影矩阵 $\mathbf{W}_i^k$ 和值投影矩阵 $\mathbf{W}_i^v$ 分别映射并分解为 $\mathbf{A}_i^k\mathbf{B}_i^k$ 和 $\mathbf{A}_i^v\mathbf{B}_i^v$ 两个低秩矩阵对。

值重建矩阵与输出投影矩阵的离线融合。基于注意力输出的计算公式,Palu在离线阶段将值重建矩阵 $\mathbf{B}_i^v$ 直接吸收到输出投影矩阵 $\mathbf{W}_i^o$ 中。具体的数学推导为:$\mathbf{a}_i\mathbf{W}_i^o = (\mathbf{p}_i\mathbf{V}_i)\mathbf{W}_i^o = (\mathbf{p}_i\mathbf{H}_i^v\mathbf{B}_i^v)\mathbf{W}_i^o = \mathbf{p}_i\mathbf{H}_i^v(\mathbf{B}_i^v\mathbf{W}_i^o)$。

离线融合提升计算效率。这种离线融合设计的优势在于,它允许Palu在推理时完全跳过显式重建完整值向量的步骤。这一优化不仅减少了矩阵乘法的总次数,还显著提升了整体的计算效率。

键重建矩阵与查询投影矩阵的离线融合。对于注意力分数的计算,Palu采用了类似的优化策略。矩阵 $\mathbf{B}_i^k$ 可以在离线状态下被融合到查询投影矩阵 $\mathbf{W}_i^q$ 中,其推导过程为:$\mathbf{q}_i\mathbf{K}_i^T = \mathbf{q}_i(\mathbf{H}_i^k\mathbf{B}_i^k)^T = \mathbf{x}_t\mathbf{W}_i^q(\mathbf{B}_i^k)^T(\mathbf{H}_i^k)^T = \mathbf{x}_t\Big(\mathbf{W}_i^q(\mathbf{B}_i^k)^T\Big)(\mathbf{H}_i^k)^T$。由于 $\mathbf{B}_i^k \in \mathbb{R}^{r \times d_h}$ 且 $\mathbf{W}_i^q \in \mathbb{R}^{d \times d_h}$,融合后的矩阵 $(\mathbf{W}_i^q(\mathbf{B}_i^k)^T)$ 的维度变为了 $\mathbb{R}^{d \times r}$。这种融合通过在计算注意力分数时降低矩阵维度,极大地提升了计算效率。

RoPE对注意力分数融合的限制及动态重建。对于最近的LLMs(如Llama系列),它们在查询和键状态相乘之前应用了旋转位置嵌入(RoPE)【索引编号:[51],Roformer: Enhanced transformer with rotary position embedding+2021+CoRR+https://arxiv.org/abs/2104.09864】。由于这些位置嵌入具有非线性特征,这使得前文描述的注意力分数矩阵融合失效。为了解决这个问题,Palu在解码期间会从潜在表示中动态地重建键。作者通过定制设计的GPU内核进一步提升了这种动态重建的效率 。

非RoPE注意力机制的完美兼容。需要注意的是,对于某些位置嵌入方法(如ALiBi 【索引编号:[42],Train short, test long: Attention with linear biases enables input length extrapolation+2022+ICLR+https://openreview.net/forum?id=R8sQPpGCv0】)或新型注意力机制( 如MLA 【索引编号:[12],Deepseek-v2: A strong, economical, and efficient mixture-of-experts language model+2024+无+无】),位置嵌入并不直接应用于键状态。因此,前文描述的矩阵融合依然有效。对于这些非RoPE的注意力模块,Palu能够实现比RoPE注意力更高的加速比,因为它们通过矩阵融合完全避免了重建开销。

多头低秩分解(M-LRD)的精度衰减问题。作者将前述的按头分解方案命名为多头低秩分解(M-LRD)。研究发现,M-LRD通常会导致不可忽视的精度下降。这可能是因为SVD未能捕获跨不同注意力头共享的公共信息。因此,需要寻找替代方法来保持模型精度。

联合头低秩分解(J-LRD)的定义。一种替代方法是将所有头的权重矩阵联合起来进行分解。通过考虑组合后的权重矩阵 $\mathbf{W}_{\mathrm{joint}} = [\mathbf{W}_1, \mathbf{W}_2, \ldots, \mathbf{W}_n] \in \mathbb{R}^{d \times (d_h \cdot n_h)}$,可以执行单次低秩分解 $\mathbf{W}_{\mathrm{joint}} \approx \mathbf{A}_{\mathrm{joint}}\mathbf{B}_{\mathrm{joint}}$,其中 $\mathbf{A}_{\mathrm{joint}} \in \mathbb{R}^{d \times r_{\mathrm{joint}}}$ 且 $\mathbf{B}_{\mathrm{joint}} \in \mathbb{R}^{r_{\mathrm{joint}} \times (d_h \cdot n_h)}$。作者将此方案称为联合头低秩分解(J-LRD)。

J-LRD保持精度的优势。J-LRD的优势在于能够保留不同头之间共享的公共主成分。这是因为当SVD应用于更大的组合矩阵时,它在捕获主导成分方面特别有效,从而产生更准确的近似结果。

J-LRD的计算与重建过程。在J-LRD中,所有头共享的联合潜在表示可以通过 $\mathbf{h}_{\mathrm{joint}} = \mathbf{x}\mathbf{A}_{\mathrm{joint}}$ 计算得出。在解码期间,每个头的原始状态可以通过公式 $[\mathbf{y}_1, \ldots, \mathbf{y}_n] = \mathbf{h}_{\mathrm{joint}}\mathbf{B}_{\mathrm{joint}}$ 进行重建。

J-LRD的高推理开销。尽管J-LRD能更好地保留模型精度,但它在解码期间引入了显著的计算和内存开销。具体而言,重建一个头的键或值状态所需的总浮点运算次数(FLOPs)变为了 $r_{\mathrm{joint}} \cdot d_h \cdot n$。假设低秩潜在表示的总大小相同(即 $r_{\mathrm{joint}} = \sum_{i=1}^n r_i$),其重建成本是M-LRD(总FLOPs为 $r_i \cdot d_h \cdot n$)的 $n$ 倍。此外,考虑到矩阵融合,J-LRD的融合矩阵大小为 $r_{\mathrm{joint}} \cdot d \cdot n$,同样是M-LRD的 $n$ 倍,导致了大幅增加的内存消耗。

组头低秩分解(G-LRD)的提出。为了在精度和重建成本之间取得平衡,作者提出了组头低秩分解(G-LRD)。G-LRD将一组头的矩阵结合在一起进行分解。通过组合权重矩阵,它在限制计算开销的同时,捕获了每个组内的共享信息,从而保持了精度。

G-LRD的分解与重建计算。为了说明G-LRD的过程,考虑一组 $s$ 个头的权重矩阵 $\mathbf{W}_{g_j} = [\mathbf{W}_{j,1} \ldots \mathbf{W}_{j,s}]$,其中 $\mathbf{W}_{g_j} \in \mathbb{R}^{d \times (d_h \cdot s)}$。对其进行低秩分解 $\mathbf{W}_{g_j} \approx \mathbf{A}_{g_j}\mathbf{B}_{g_j}$,其中 $\mathbf{A}_{g_j} \in \mathbb{R}^{d \times r_g}$ 且 $\mathbf{B}_{g_j} \in \mathbb{R}^{r_g \times (d_h \cdot s)}$。同一组内注意力头共享的潜在表示计算为 $\mathbf{h}_{g_j} = \mathbf{x}\mathbf{A}_{g_j}$。在解码期间,每个头的原始键或值可以通过 $[\mathbf{y}_{j,1} \ldots \mathbf{y}_{j,s}] = \mathbf{h}_{g_j}\mathbf{B}_{g_j}$ 重建。

G-LRD的开销分析。在G-LRD中,重建每个头的键和值的FLOPs为 $r_g \cdot d_h \cdot n_g$,其中 $n_g = \frac{n}{s}$ 是组的数量。与J-LRD相比,假设总秩大小相同($r_g \cdot n_g = r_{\mathrm{joint}}$),G-LRD将重建成本降低了 $n_g$ 倍。同样,G-LRD也将融合矩阵的大小减小了 $n_g$ 倍。总而言之,G-LRD在计算开销和近似精度之间提供了一个折中方案。

在不同粒度下执行分解。联合分解多个头可以获得更高的精度。假设潜在表示的总大小相同(即 $4 \cdot r_i = 2 \cdot r_g = r_{\mathrm{joint}}$),联合头分解方案中用于重建开销的 FLOPs 是多头分解方案的 4 倍。
在不同粒度下执行分解。联合分解多个头可以获得更高的精度。假设潜在表示的总大小相同(即 $4 \cdot r_i = 2 \cdot r_g = r_{\mathrm{joint}}$),联合头分解方案中用于重建开销的 FLOPs 是多头分解方案的 4 倍。

利用Fisher信息评估矩阵重要性。为了将理想的秩大小分配给分解目标,准确估计目标矩阵(例如分组权重)的重要性至关重要。在Palu中,作者确定Fisher信息【索引编号:[34],A tutorial on fisher information+2017+无+无】【索引编号:[31],Group fisher pruning for practical network compression+2021+ICML+http://proceedings.mlr.press/v139/liu21ab.html】为一个准确的近似器,因为它能够量化每个参数的信息量。随后,作者采用Fisher信息的总和来估计每个线性层权重矩阵的重要性【索引编号:[1] ,Zero-cost proxies for lightweight NAS+2021+无+https://openreview.net/forum?id=0cmMMy8J5q】 。

按Fisher信息比例分配秩大小。假设压缩敏感度与Fisher信息成正比,作者通过计算每个权重矩阵的Fisher信息与所有分解目标总Fisher信息的比率来确定其秩。使用这个比率来分配压缩率(即秩级别 $r$),确保更重要的层保留更高的秩级别。

低秩分解引入严重异常值。作者将量化技术集成到Palu中以进一步压缩KV-Cache。研究观察到,低秩压缩后的潜在表示存在严重的异常值,这限制了量化在Palu中的适用性。与之前KV-Cache量化文献【索引编号:[33],Kivi: A tuning-free asymmetric 2bit quantization for kv cache+2024+arXiv+https://arxiv.org/abs/2402.02750】【索引编号:[24] ,Kvquant: Towards 10 million context length llm inference with kv cache quantization+2024+arXiv+https://arxiv.org/abs/2401.18079】中描述的自然异常值不同,这些异常值是由基于SVD的低秩分解诱发的 。

异常值分布规律及其对量化的破坏。Fig 4(a) 展示了来自Llama-2模型某一层使用G-LRD的低秩压缩键状态的分布情况。重复的异常值模式出现在每个分解组的开头,这是因为SVD将较大的特征值排列在初始行或列中,导致潜在表示中的值迅速下降。这种模式拉伸了数据分布,严重损害了量化精度。

Llama-2 第 4 个注意力层低秩键缓存的激活分布。
Llama-2 第 4 个注意力层低秩键缓存的激活分布。

利用Hadamard变换消除异常值。受近期LLM量化文献【索引编号:[4],Quarot: Outlier-free 4-bit inference in rotated llms+2024+arXiv+https://arxiv.org/abs/2404.00456】【索引编号:[55] ,Quip#: Even better llm quantization with hadamard incoherence and lattice codebooks+2024+arXiv+https://arxiv.org/abs/2402.04396】的启发,作者应用了Walsh-Hadamard变换(WHT)【索引编号:[16] ,Unified matrix treatment of the fast walsh-hadamard transform+1976+IEEE Transactions on Computers+无】来消除异常值(如Fig 4(b)所示),从而实现了高量化精度。

Hadamard矩阵的无开销离线融合。然而,这种变换引入了额外的矩阵乘法及相关的运行时开销。与之前必须在量化KV-Cache时应用在线WHT的方法不同,作者通过将Hadamard矩阵集成到低秩分解权重中优化了这一过程,计算公式为 $\mathbf{W} \approx \mathbf{A}\mathbf{B} = (\mathbf{A}\mathbf{R})(\mathbf{R}^T\mathbf{B}) = \hat{\mathbf{A}}\hat{\mathbf{B}}$,其中 $\mathbf{R}$ 是Hadamard矩阵。这种优化允许Palu将提出的低秩压缩技术与低比特量化无缝结合,且不产生额外的计算开销。

实验环境

实验结果

结论

本文提出了 Palu,一个新颖的 KV-Cache 压缩框架,通过分解线性投影权重矩阵并缓存压缩后的潜在表示来减少内存。Palu 引入了多项优化,包括组头低秩分解、自动秩分配算法、量化兼容性增强以及带有算子融合的定制内核。实验证明,在 50% 的低秩压缩和 4-bit 量化下,Palu 能够将 RoPE 注意力模块加速高达 2.91 倍,端到端加速高达 2.2 倍,同时在各大基准测试中保持了强大的模型精度。未来的工作可以探索对分解后的权重矩阵进行进一步量化,并利用低精度硬件(如 INT4 Tensor Cores)来进一步优化在线重建的效率。

补充细节

SVD在LLM压缩中的应用。多项先前的研究探索了使用SVD压缩LLM。早期的工作(Noach & Goldberg, 2020)将标准SVD直接应用于权重矩阵,导致了显著的压缩误差。FWSVD使用Fisher信息来优先处理重要参数,而ASVD则考虑了激活的异常值。SVDLLM进一步最小化了每个奇异值的压缩损失。与这些压缩模型权重的方法不同,Palu专门专注于减小KV-Cache的大小。一项同期工作(Yu et al., 2024)也探索了使用低秩投影压缩KV-Cache,但该方法需要在分解后进行LoRA微调。相比之下,Palu直接分解权重矩阵,无需微调即可保持精度,并引入了秩搜索、量化集成和优化GPU内核等额外创新。

KV-Cache压缩。量化是压缩KV-Cache的广泛使用技术。Atom应用简单的逐词元量化,WKVQuant引入了两级方案以提高精度。KIVI对键和值使用逐通道和逐词元量化,并结合组大小为32的细粒度组量化。KVQuant采用类似设置,但结合了非均匀量化和稀疏矩阵来处理异常值。在这些方法之上,GEAR添加了一个低秩矩阵来补偿量化误差。在Palu中,作者利用低秩技术正交地开发隐藏维度的冗余,仅使用简单的逐词元量化就取得了出色的压缩结果。

MLA机制对比。最近发布的DeepSeek-V2模型【索引编号:[12],Deepseek-v2: A strong, economical, and efficient mixture-of-experts language model+2024+无+无】引入了MLA机制,该机制通过将键和值向下投影到低秩空间并在运行时将其重建为满秩来减小KV-Cache大小。尽管MLA在宏观上(特别是使用J-LRD时)似乎与Palu相似,但设计和推导过程有着根本的不同。MLA是一种需要预训练的新型注意力机制,而Palu是专门为后训练集成设计的,专注于转换现有的MHA或GQA模型以支持低秩压缩的KV-Cache。

支持在线重建的注意力分数计算内核。Palu的核心思想是利用低秩潜在表示来减少数据传输开销,从而加速注意力机制。系统存储并传输压缩后的低秩潜在表示 $\mathbf{H} \in \mathbb{R}^{L \times r}$,而不是完整的键矩阵。在计算期间,定制的GPU内核使用重建矩阵 $\mathbf{B} \in \mathbb{R}^{r \times d_h}$ 执行动态重建,生成恢复的键矩阵 $\mathbf{K} \in \mathbb{R}^{L \times d_h}$。随后,查询向量 $\mathbf{q} \in \mathbb{R}^{1 \times d_h}$ 与重建的键相乘获得注意力分数。为了高效利用并行性,内核沿序列长度维度 $L$ 执行分块(Tiling)。序列被拆分为大小为 $L_{\mathrm{tile}}$ 的较小块,每个线程块独立地从 $\mathbf{H}$ 重建子矩阵 $\mathbf{H}_i$,应用RoPE位置嵌入,最后执行矩阵-向量乘法。这种设计确保所有中间计算完全保留在片上共享内存中,从而最小化高延迟内存访问并实现显著加速。

用于计算具有在线重建的注意力分数的融合 GPU 内核的图示。在该图中,$\mathbf{q}$ 代表查询向量,$\mathbf{H}$ 代表低秩压缩的键状态,$\mathbf{B}$ 代表重建矩阵。
用于计算具有在线重建的注意力分数的融合 GPU 内核的图示。在该图中,$\mathbf{q}$ 代表查询向量,$\mathbf{H}$ 代表低秩压缩的键状态,$\mathbf{B}$ 代表重建矩阵。

分解粒度对权重存储的影响。评估KV-Cache压缩不仅要看压缩率,还要考虑整体内存节省。J-LRD分解方案由于联合分解所有头的投影,会导致权重尺寸增加(例如增加约40%的存储)。相比之下,M-LRD和优化的G-LRD方案涉及非方形目标矩阵。例如,在组大小为4的G-LRD中,Llama-2-7b模型的拼接矩阵大小为 $4096 \times 512$。计算表明,这不仅没有增加额外存储成本,反而实现了21.25%的额外内存节省。此外,K和V投影的权重仅占Transformer块中7个线性层的2个(约占Llama-2-7b参数的16%),限制了对整体内存的负面影响。

整体内存占用分析。通过分析Llama-2-7B在不同序列长度下的总内存使用量(包含模型权重和KV-Cache),发现在64K序列长度时,KV-Cache占用了78%的总内存。在50%低秩压缩下,Palu有效地将总内存使用量降低了1.7倍;当结合2-bit量化时,总内存使用量进一步降低了4.6倍。

Llama-2-7B 中各种序列长度的总内存使用情况。对于 Palu,低秩压缩率为 50%。
Llama-2-7B 中各种序列长度的总内存使用情况。对于 Palu,低秩压缩率为 50%。

结合LoRA微调恢复精度。LoRA可以作为一种压缩后恢复技术,用于恢复压缩后丢失的信息。在Palu中,引入了额外的低秩矩阵 $\mathbf{A}_{r'}' \in \mathbb{R}^{d \times r'}$ 和 $\mathbf{B}_{r'}' \in \mathbb{R}^{r' \times d}$ 来细化原始低秩投影:$\mathbf{h} = \mathbf{A}\mathbf{x} + \mathbf{A}_{r'}'\mathbf{B}_{r'}'\mathbf{x}$。其中 $\mathbf{A}$ 和 $\mathbf{B}$ 是固定的分解参数,而带撇号的矩阵是可训练参数。实验结果表明,结合LoRA后,J-LRD的平均性能下降仅为1.00%。G-LRD (组大小为4) 和 M-LRD 的平均下降分别改善至 2.01% 和 5.14%。特别是 G-LRD 在结合 LoRA 后,与 J-LRD 的精度差距缩小至仅 1.03%。

13B模型的零样本精度。在 Llama-2-13B 模型上以 50% 压缩率进行的评估显示,无论是使用 J-LRD、G-LRD 还是 M-LRD,Palu 都能实现具有竞争力的精度下降(约 3% 或更少)。这意味着用户可以优先采用 M-LRD 来进一步优化效率。

组大小的影响。消融实验(表6)表明,随着组大小的增加,共享信息的数量也随之增加,从而带来了性能的提升。在 Llama-2-7B 的 50% 压缩率下,组大小从 1 (M-LRD) 增加到 32 (J-LRD),困惑度从 6.81 稳步下降至 5.62。

Hadamard变换的影响。消融实验(表7)证实了Walsh-Hadamard变换(WHT)的优势。在3-bit量化级别,Hadamard变换仅带来轻微的困惑度改善。然而,在更极端的2-bit量化下,Hadamard变换带来了显著的 4.17 困惑度提升(从 10.58 降至 6.41)。并且由于Palu通过离线预处理优化了WHT过程,这不会在推理期间带来额外开销。

自动秩分配与均匀分配对比。消融实验(表8)表明,应用秩搜索带来了显著的性能提升。在50%压缩率下,困惑度显著降低了1.36;在70%压缩率下降低了2.18。对Llama-2-7B层级低秩压缩率的可视化表明,分配结果是非均匀的。具体而言,Value 投影通常比 Key 被分配更高的秩;同时,前半部分层被分配了更高的秩,表明它们在保持模型性能方面更为重要。

总体压缩率为 50% 的 Llama-2-7B 上逐层低秩压缩率的可视化。此处,使用提出的基于 Fisher 信息的自动秩分配算法来分配压缩率(即秩)。
总体压缩率为 50% 的 Llama-2-7B 上逐层低秩压缩率的可视化。此处,使用提出的基于 Fisher 信息的自动秩分配算法来分配压缩率(即秩)。