MagicPIG: LSH Sampling for Efficient LLM Generation
MagicPIG: LSH Sampling for Efficient LLM Generation
发表时间: 2025-04 · arXiv:2410.16179 (ICLR 2025)
原文: https://arxiv.org/abs/2410.16179
速读
一句话结论 本文提出了基于局部敏感哈希(LSH)采样的异构推理系统 MagicPIG,通过将大模型的 KV Cache 和注意力计算卸载到 CPU,并用无偏采样替代传统的 TopK 稀疏注意力,在极低计算开销下保持了长上下文任务的高精度,并将解码吞吐量提升了最高 5 倍。
要解决什么问题 长上下文大模型在自回归生成时,KV Cache 会随批次大小和序列长度线性增长,导致严重的显存瓶颈和访存受限(Memory-bound),使得 GPU 算力利用率极低。为了缓解这个问题,现有方法通常采用动态稀疏或基于 TopK 的注意力近似机制,只计算得分最高的少数 Token。但这卡在了三个致命机制上:首先,注意力分布并不总是高度稀疏的,尤其是在需要综合全局信息的聚合任务(如词汇提取)中,注意力呈现长尾分布,强制截断的 TopK 会引入严重偏差,导致模型能力断崖式掉点;其次,看似极高的稀疏度往往是由“注意力下沉(Attention Sink)”现象制造的假象,除了初始 Token 外,其余 Token 的注意力分布相对均匀,且 Query 和 Key 在几何空间上分布于方向几乎相反的狭窄锥体中,这导致精确搜索 TopK 的开销极大;最后,现有的稀疏方法大多只能减少计算量,却无法真正缩减 KV Cache 占用的总显存,在显存紧张时依然无法扩大上下文长度或批次大小。
怎么做的 核心思路是放弃有偏的 TopK 截断,将注意力近似转化为一个带有理论保证的无偏采样估计问题。作者发现,按注意力分数的真实分布进行采样,能以极小的计算预算大幅降低估计误差。为了在未知完整注意力分数的前提下实现高效采样,MagicPIG 采用了自归一化重要性采样(Self-normalized Importance Sampling),并利用局部敏感哈希(LSH)来近似采样概率。具体而言,方法包含三个关键设计。第一是基于 LSH 的采样估计:利用 SimHash 算法,通过 $K$ 个随机投影哈希函数和 $L$ 个哈希表,让与 Query 余弦相似度高的 Key 有更大概率发生哈希碰撞并被采样。采样概率 $u_i$ 与碰撞机制定义为:
其中 $p_i = 1 - \frac{1}{\pi} \arccos \frac{q k_i^T}{|q| \cdot |k_i|}$。最终的注意力输出通过采样集合 $S$ 进行重要性加权估计:
效果如何 实验覆盖了 Llama-2-7b、Llama-3.1-8B/70B 以及 CodeLlama 等模型,在 A100、L20 和单张 RTX 4090 上进行了测试。评估任务包括 lm-eval-harness 的中等长度任务、LongBench 和 RULER(上下文最高达 256K)的长文本任务。对比的基线方法主要有两个:代表性能上限的 GPU 完整注意力机制(Full Attention),以及代表动态稀疏路线的强基线方法 Quest。量化结果显示,MagicPIG 在仅使用完整注意力 2% 到 5% 计算开销的情况下,在各类任务上的精度损失均不到 2%。在信息聚合类任务上,它的表现甚至超越了精确的 TopK 注意力。在系统效率方面,由于成功将 KV Cache 卸载到 CPU,MagicPIG 能够支持比基线大 12 倍以上的批次大小,从而在不同硬件上将解码吞吐量提升了 1.5 倍到 5 倍。在单张 24GB 显存的 RTX 4090 上,运行 96K 上下文的 Llama-3.1-8B 模型时,单请求解码延迟仅为 54 毫秒。作者也承认了该方法的局限性:它强依赖于 CPU 拥有充足的 DRAM 来存储卸载的 KV Cache 和庞大的哈希表,不适用于内存受限的场景;此外,目前 MagicPIG 仅针对解码阶段进行了优化,尚未在预填充阶段实现该机制。
Zhuoming Chen, Ranajoy Sadhukhan, Zihao Ye, Yang Zhou, Jianyu Zhang, Niklas Nolte, Yuandong Tian, Matthijs Douze, Leon Bottou, Zhihao Jia, Beidi Chen / Carnegie Mellon University, University of Washington, New York University, Meta AI
主要贡献
具有长上下文窗口的大型语言模型(LLMs)受到了广泛关注,但在自回归生成过程中,为了避免重复计算而存储的键值(KV)缓存成为了一个独特的性能瓶颈。KV缓存随着批处理大小和序列长度呈线性增长,占据了大量GPU内存,导致LLM生成极度受限于内存带宽,从而使得GPU计算能力利用率低下。为了利用“注意力具有稀疏性”这一普遍共识,以往的研究提出了各种动态稀疏或基于TopK的注意力近似方法。
本文的核心研究目标是提出一种理想的稀疏注意力近似方法,该方法需在提供理论保证的同时,在多样化的下游任务中保持全精度、具有低成本的KV缓存选择开销,并能有效节省GPU内存。
本文的主要创新点如下:
1. 揭示TopK注意力的局限性:首次证明了TopK注意力本身在某些需要聚合全上下文的下游任务中会导致严重的质量下降,因为注意力并不总是如预期般稀疏。
2. 提出基于采样的注意力估计:证明了相较于选择具有最高注意力分数的键和值,具有理论保证的采样方法能够为注意力输出提供更好的无偏估计。
3. 设计MagicPIG系统:提出了一种基于局部敏感哈希(Locality Sensitive Hashing, LSH)的异构系统——MagicPIG。该系统将LSH哈希表存储在CPU上,并在CPU上运行注意力计算,从而大幅减少了注意力计算的工作量。
4. 突破硬件限制:通过CPU-GPU协同设计,MagicPIG能够支持更长的上下文和更大的批处理大小。在单张RTX 4090上,针对96k tokens上下文的Llama-3.1-8B-Instruct模型,MagicPIG可实现54ms的解码延迟,并在各种GPU硬件上将解码吞吐量提升高达5倍。
背景知识与关键观察
注意力估计问题定义。在LLM解码阶段,自注意力机制通过 $o = \mathrm{Softmax}(\frac{qK^T}{\sqrt{d}})V = wV$ 计算先前值的加权平均。其中 $d$ 为头部维度,$n$ 为上下文大小。目标是找到采样矩阵 $\Pi \in \mathbb{R}^{n \times m}$ 和对角矩阵 $D \in \mathbb{R}^{m \times m}$,以最小化误差 $\delta = ||wV - w\Pi D\Pi^T V||$,且计算预算 $m \ll n$。
局部敏感哈希(LSH)基础。LSH通过将相似输入分配相同哈希码的概率最大化来进行近似最近邻搜索。SimHash【[13], Similarity estimation techniques from rounding algorithms + 2002 + STOC】是基于余弦相似度的LSH变体。对于向量 $x \in \mathbb{R}^d$,SimHash生成随机超平面 $w$ 并返回 $\mathrm{Sign}(w^Tx)$。两个向量 $x, y$ 共享相同符号的概率为 $p = 1 - \frac{\theta}{\pi}$,其中 $\theta = \arccos \frac{x^Ty}{||x||\cdot||y||}$。
TopK注意力的致命弱点。尽管基于搜索的稀疏注意力算法【[63], Quest: Query-aware sparsity for efficient long-context llm inference + 2024 + arXiv】试图逼近TopK,但TopK注意力本身是有偏且不准确的。特别是当注意力分数分布呈现长尾特性且计算预算有限时,TopK在需要利用完整上下文的聚合任务中表现极差。实验表明,前20%的tokens仅能覆盖70%~80%的注意力分数,导致不可忽视的15%~20%的估计误差。
注意力的几何特征观察。通过研究查询 $q$ 和键 $k$ 的几何结构,得出以下关键观察:
1. 初始token(即注意力Sink,记为 $k_{sink}$)的键状态对于任意输入几乎保持不变,相互间的余弦相似度大于0.99【[68], Efficient streaming language models with attention sinks + 2023 + arXiv】。
2. 键状态的中心(即平均键 $k_{avg} = \frac{1}{n}\sum_{i=1}^n k_i$)在不同输入句子中方向保持稳定,相似度超过0.9。
3. $k_{avg}$ 和 $k_{sink}$ 的方向几乎完全相反,余弦相似度在 $-0.9 \sim -0.8$ 之间。
这些几何特征表明,注意力Sink独立于输入产生高稀疏性,而其他部分分布更均匀。简单的TopK会过度赋予Sink权重,从而丢失上下文信息。此外,$q$ 和 $k$ 的不对齐也导致了搜索困难【[44], Retrievalattention: Accelerating long-context llm inference via vector retrieval + 2024 + arXiv】。
通过采样估计注意力。将注意力输出 $o$ 重写为来自分布 $w$ 的期望值,即 $o = \mathbb{E}_{i \sim w}(v_i)$。定义3.1(Oracle采样估计):给定采样预算 $\mathcal{B}$ 和归一化注意力分数 $w$,从 $w$ 中独立采样 $\mathcal{B}$ 个元素,输出估计为 $\bar{o} = \frac{1}{\mathcal{B}}\sum_{j=1}^{\mathcal{B}} v_{i_j}$。
定理3.2:Oracle采样估计是无偏的,且协方差的迹随 $\mathcal{B}$ 单调递减。
定理3.3:预期计算预算 $\mathbb{E}(|S|)$ 具有上限 $1 + \mathcal{B}\epsilon$,其中 $\epsilon = 1 - \max_i w_i$。这证明了实际计算成本通常远小于采样预算。
方法细节
自归一化重要性采样估计。由于获取精确分布 $w$ 需要计算所有的 $qk_i^T$,Oracle采样无法带来实质加速。因此,引入自归一化重要性采样,通过从提议分布 $u$ 中采样索引 $i_1, i_2, ..., i_{\mathcal{B}}$ 来估计未知分布 $w$。由此产生的估算器为 $X^{\mathrm{IS}} = \frac{1}{\widetilde{Z}} \sum_{j=1}^{\mathcal{B}} \frac{\widetilde{w_{i_j}}}{u_{i_j}} v_{i_j}$,其中 $\widetilde{Z} = \sum_{j=1}^{\mathcal{B}} \frac{\widetilde{w_{i_j}}}{u_{i_j}}$,且 $\widetilde{w_i} = e^{\frac{qk_i^T}{\sqrt{d}}}$。该估算器具有 $\mathbb{P}[\lim_{\mathcal{B} \to \infty} X^{\mathrm{IS}} = o] = 1$ 的优良性质。为了最小化方差,提议分布 $u$ 必须满足 $u_i \propto \widetilde{w_i}|v_i - o|$。
利用LSH进行方差缩减。对目标分布进行分解:$\widetilde{w_i}|v_i - o| = \exp(\frac{qk_i^T}{\sqrt{d}} + \log|v_i - o|)$。经验观察表明,$\log|v_i - o|$ 的波动相比 $\frac{qk_i^T}{\sqrt{d}}$ 并不显著,因此将 $u$ 的要求简化为与 $qk_i^T$ 共享相同的峰值。接着,通过向量变换 $r = \max_{1 \le i \le n} |k_i|$,$\bar{q} = [q, 0]$,$\bar{k_i} = [k_i, \sqrt{r^2 - |k_i|^2}]$,将内积转换为余弦相似度。基于这一转换,利用SimHash【[56], Simhash: Hash-based similarity detection + 2007】构建概率分布 $u_i = \mathbb{P}[h(q) = h(k_i)]$,该概率与 $\cos \frac{qk_i^T}{|q| \cdot |k_i|}$ 单调相关。
估算器近似与哈希函数选择。由于哈希提供的概率未归一化,MagicPIG对估算器进行了调整,计算 $X = \frac{\sum_{i \in S} \frac{\widetilde{w_i}}{u_i} v_i}{\sum_{i \in S} \frac{\widetilde{w_i}}{u_i}}$。在哈希函数选择上,利用 $K \times L$ 个随机向量进行SimHash。对于 $L$ 个哈希表中的每一个,保留投影的符号生成 $K$ 位哈希值。仅当键 $k_i$ 在至少两个哈希表中与 $q$ 共享哈希值时,才对其进行采样。对应的采样概率为 $u_i = 1 - (1 - p^K)^L - Lp^K(1 - p^K)^{L-1}$,其中 $p_i = 1 - \frac{1}{\pi} \arccos \frac{qk_i^T}{|q| \cdot |k_i|}$。
数据预处理与居中化。由于键几乎总是集中在查询的一侧(除了初始token),随机投影无法有效区分键,会导致均匀的采样概率。因此,在构建哈希表前,MagicPIG先对 $k_i$ 向量进行中心化处理($\bar{k_i} = k_i - \frac{1}{n}\sum_{i=1}^n k_i$),使键分布更优且保持计算上的等价性。
算法执行过程。结合上述设计,MagicPIG的执行步骤如下:
首先,计算新查询 $\pmb{q}$ 的哈希码 $\pmb{q}_{\mathrm{code}} = \mathbf{Encode}(\pmb{q}, \pmb{W})$。
再者,查询哈希表以获取采样集合 $S = \mathbf{Query}(HT, \pmb{q}_{\mathrm{code}})$,并提取对应的键值 $\pmb{K}_S = K[S], \pmb{V}_S = V[S]$。
接着,计算 $\pmb{q}$ 与采样 $\pmb{K}_S$ 及静态缓存 $\pmb{K}_T$ 的内积:$\pmb{w}_S = \pmb{q}\pmb{K}_S^T$,$\pmb{w}_T = \pmb{q}\pmb{K}_T^T$。
然后,计算碰撞概率 $\pmb{p} = 1 - \frac{1}{\pi} \arccos(\pmb{w}_S / (||\pmb{q}|| \cdot ||\pmb{K}_S||))$,并由此计算采样概率 $\pmb{u} = 1 - (1 - \pmb{p}^K)^L - L\pmb{p}^K(1 - \pmb{p}^K)^{L-1}$。
最后,结合静态缓存 $T$(如Sink token)计算注意力输出估计:$\bar{o} = \mathbf{Softmax}(\frac{[\pmb{w}_S, \pmb{w}_T]}{\sqrt{d}} - \log([\pmb{u}, \mathbf{1}_t]))[\pmb{V}_S, \pmb{V}_T]$。
Input: K, V, q, random projectors W, hash tables HT, static KV cache K_T, V_T
# Compute hash code for new query
q_code = Encode(q, W)
# Query hash tables to sample S
S = Query(HT, q_code)
K_S = K[S], V_S = V[S]
# Compute inner product for q and sampled K
w_S = q * K_S.T
w_T = q * K_T.T
# Compute collision probability
p = 1 - (1 / pi) * arccos(w_S / (||q|| * ||K_S||))
# Compute sampling probability
u = 1 - (1 - p**K)**L - L * p**K * (1 - p**K)**(L - 1)
# Compute attention output estimation
o_bar = Softmax([w_S, w_T] / sqrt(d) - log([u, 1_t])) * [V_S, V_T]
Return o_bar
系统协同设计。为了突破GPU VRAM的内存瓶颈,MagicPIG利用带宽为 $100-200\mathrm{GB/s}$ 的CPU DRAM作为GPU的聚合内存【[31], Fastdecode: High-throughput gpu-efficient llm serving using heterogeneous pipelines + 2024 + arXiv】。解码工作负载被划分为四个部分:
1. 参数计算:所有线性投影(MLP, $W_Q, W_K, W_V, W_O$)在GPU上运行。
2. 随机投影:为每个 $q$ 进行 $K \times L$ 次随机投影以获取哈希码。由于所有头共享相同的随机投影器,内存开销极小(400KB),受计算限制,因此放置在GPU上。
3. 检索:在CPU上的 $L$ 个预建哈希表中查找 $q$ 的哈希码。哈希表占用大量内存,适合放在CPU。
4. 注意力计算:注意力核心计算 $o = \mathrm{Softmax}(\dots)V$ 在CPU上运行。
设备端缓存(On-device cache):由于Sink token(前几个token)和局部token极有可能被采样,为进一步减少CPU工作量,MagicPIG将这些token存储在GPU上,不对其应用LSH采样,并利用递归注意力技术合并CPU和GPU的输出。
实验环境
-
数据集:
- lm-eval-harness(中等上下文):GSM8K-CoT, MMLU-Flan-Cot-Fewshot, COQA。
- LongBench(长上下文):QASPER, LCC, Repobench-P, TriviaQA, PRE, TREC。
- RULER(合成任务):13个合成任务(每任务50个示例)。
- infini_igsm(数学推理任务)。
-
模型架构:Llama-2-7b-chat, Llama-3.1-8B-Instruct, Llama-3.1-70B-Instruct, Code-Llama-13b-16K, Code-Llama-34b-16K, MegaBeam-Mistral-7B-512K, Llama3-8B-Prolong-512K。
-
硬件配置:
- GPU:80GB A100, 48GB L20, 24GB RTX 4090。
- CPU:Intel Platinum 8480+ (配合A100), Intel 8563C (配合L20)。
-
软件配置:GPU部分使用原生PyTorch实现,CPU部分使用FBGEMM实现,采用bfloat16精度。基线方法包括全注意力(Full Attention)和Quest【[63], Quest: Query-aware sparsity for efficient long-context llm inference + 2024 + arXiv】。
实验结果
1. 准确率保持评估
- 实验内容:在lm-eval-harness、LongBench和RULER上对比MagicPIG、全注意力和Quest的准确率。
- 实验结果:MagicPIG在所有任务中均保持了极高的准确率(性能下降不到2%)。在lm-eval-harness中,MagicPIG(10,220)配置下Llama-2-7b的平均得分为47.4,而Quest(16,0.05)仅为41.3;在LongBench中,MagicPIG在仅消耗2%~5%计算成本的情况下,平均得分达到70.3~70.7,媲美全注意力的71.2。在RULER长达96K的上下文中,MagicPIG同样展现出极强的鲁棒性。
- 分析结论:由于引入了LSH采样,MagicPIG的搜索/采样开销($\mathrm{Cost}_1$)比Quest低一个数量级,能够在仅使用基线一半计算成本的情况下实现同等或更高的准确率。
2. 硬件效率与系统性能评估
- 实验内容:在A100 (34B模型, 16K上下文)、L20 (13B模型, 16K上下文) 和 RTX 4090 (8B模型, 96K上下文) 三种场景下评估解码吞吐量和延迟。
- 实验结果:MagicPIG显著提高了所有场景的解码吞吐量(A100提升1.5倍,L20提升5.0倍,RTX 4090提升3.3倍)。在RTX 4090上服务96K上下文的单请求生成时,实现了54ms的极低延迟。
- 分析结论:得益于将KV缓存卸载至CPU,MagicPIG能够容纳比GPU全注意力基线大得多的批处理大小(超过12倍),这是吞吐量大幅提升的核心原因。
3. 消融实验
- 实验内容:验证中心化(Centering)操作的必要性,并对比MagicPIG与确切TopK在聚合任务中的表现。
- 实验结果:如Figure 9a所示,如果不进行中心化,检索任务(NIAH)的准确率降至几乎为零,FWE任务降至65%。在CWE和FWE两个聚合任务中,基于采样的MagicPIG分别击败了确切的TopK注意力达3%和8%的幅度。
- 分析结论:中心化对于LSH采样至关重要,因为键的方向几乎相反,不中心化会导致无法采样到有效的键。同时,采样机制确实超越了TopK,证明了纯搜索算法无法达到的性能上限。
结论
本文首先揭示了TopK注意力近似在解决长上下文LLM生成的计算和内存挑战时的局限性。随后证明了Oracle采样能够超越TopK,并提出了一种利用LSH采样近似Oracle采样的创新方法——MagicPIG。MagicPIG通过将哈希表和减少后的注意力计算卸载到CPU的系统协同设计,显著降低了注意力计算的工作量,同时在多样化任务中保持了高准确率。实验结果表明,MagicPIG在多种硬件配置下大幅提高了吞吐量并降低了延迟,全面优于传统的TopK注意力机制。其理论的严谨性、鲁棒性和可扩展性为注意力近似方法及算法-硬件协同设计开辟了新的机遇。未来的工作包括将MagicPIG扩展至预填充(Prefilling)阶段、应用更先进的交叉多胞形哈希(Cross-polytope hash)以减小哈希表体积,以及探索适用于高端GPU的纯GPU LSH加速方案。
附录与补充细节
A 定理证明补充。附录提供了定理3.2(Oracle采样无偏性及方差递减)和定理3.3(计算预算上限约束)的完整数学推导,利用Jensen不等式和凸函数性质严格证明了 $\mathbb{E}[|S|] \le 1 + \epsilon \mathcal{B}$。
B Oracle采样策略分析。理论上,保证最低方差的最优采样概率应为 $u'_i \propto w_i ||v_i||$。但MagicPIG未采用此策略,原因有二:首先,Sink token的值范数(Value norm)显著小于其他token(如图11所示),若按此概率采样会降低其被选中的概率,从而影响注意力的功能;其次,该策略会导致概率分布比 $w_i$ 更平缓,从而大幅增加计算成本。
D 扩展评估。
- 超长上下文:在MegaBeam-Mistral-7B-512K和Llama3-8B-Prolong-512K模型上扩展至256K上下文,MagicPIG(9,120)在256K下仍保持80.1%的准确率,远超Quest的78.5%。
- 模型扩展:在70B级别的Llama-3.1-70B-Instruct上,MagicPIG(9,120)在64K上下文中以4%的计算成本实现了89.8%的准确率(全注意力为90.3%)。
- 数学推理任务:在infini_igsm任务中,MagicPIG在所有复杂度(2-Ops, 4-Ops, 5-Ops)下均一致优于Quest,而TopK注意力在此类任务中遭遇了显著的性能退化。
E LSH超参数 (K, L) 的选择。
- 作用机制:K决定了空间划分的细粒度($2^K$ 个子空间)。K过小会导致采样过多不相关的键(增加计算成本);K过大则碰撞概率极低,需要增加L(哈希表数量)来保证采样数量,这会带来巨大的CPU内存开销。
- 内存开销:在Llama-3.1-8B-Instruct(96K上下文)中,(10, 150)配置下哈希表占用14GB内存,而(11, 300)配置则占用28GB。内存开销随上下文长度和模型头数线性增长。
- 选择策略:K是高度敏感的参数,MagicPIG通过离线消融实验确定 K=8~10 为最佳通用区间。在确定K后,通过调整L来达到目标计算预算。例如,为了将计算成本控制在5%以下且L低于200,(8, 75)、(9, 120)和(10, 150)是理想的配置组合。
F TopK与采样的直观对比。附录提供了一个动物园食物消耗估计的直观例子:TopK方法(仅选取数量最多的动物种类)会赋予高消耗动物不成比例的权重,导致估计值严重偏高(有偏估计);而基于分布概率的放回采样机制,不仅能够提供无偏的估计值,还能随着采样预算的增加有效降低方差。这通俗地解释了为什么采样在聚合任务中超越了TopK。
💬 评论讨论
欢迎在这里分享您的想法和见解!