发表时间: 2023-09 · arXiv:2309.17453 (ICLR 2024)
原文: https://arxiv.org/abs/2309.17453
Guangxuan Xiao, Yuandong Tian, Beidi Chen, Song Han, Mike Lewis
Massachusetts Institute of Technology, Meta AI, Carnegie Mellon University, NVIDIA
一句话结论 提出了一种名为 StreamingLLM 的框架,通过在缓存中永久保留最初几个 token 作为“注意力汇聚点”并结合滑动窗口,使得无需微调的现有大语言模型能够稳定处理超过 400万 token 的无限长度文本流。
要解决什么问题 在多轮对话等流式应用中部署大语言模型面临两个严重的卡点。首先,标准的密集注意力机制(Dense Attention)在解码阶段会缓存所有历史 token 的键值状态(KV),这会导致显存消耗和解码延迟随文本长度线性增长,最终引发显存溢出。其次,现有模型无法外推到超过其预训练长度的文本。为了节省显存,一种直观的做法是“窗口注意力”(Window Attention),即只在缓存中保留最近几个 token 的 KV 状态。然而这种方法存在致命缺陷:一旦文本长度超过缓存大小,导致序列最开始的 token 被挤出缓存,模型的生成质量就会瞬间崩溃。另一种替代方案是“带重计算的滑动窗口”(Sliding Window with Re-computation),它在每生成一个新 token 时都重新计算最近窗口内的 KV 状态,虽然效果好但其二次方的计算复杂度导致速度极慢。窗口注意力失效的根本机制在于注意力机制中的 SoftMax 操作。SoftMax 要求所有上下文 token 的注意力分数总和必须为一,因此即使当前的查询向量与前面的 token 没有强烈的语义匹配,模型也必须将这些“无用”的注意力值分配出去。由于自回归语言模型的顺序生成特性,序列最初的几个 token 对后续所有 token 都是可见的,这使得它们最容易被训练成吸收这些多余注意力分数的“垃圾桶”。如果这些初始 token 的 KV 状态被移出缓存,SoftMax 函数的分母就会丢失巨大的一部分,导致注意力分数分布发生剧烈偏移,进而使模型的困惑度飙升。
怎么做的 StreamingLLM 的核心思路是利用上述“注意力汇聚点”(Attention Sink)现象,在不进行任何模型微调的情况下,强行在缓存中永久保留序列最初几个 token 的 KV 状态,以此来锚定注意力计算,维持正常的注意力分数分布。该方法的关键设计将 KV 缓存分成了两个独立的部分:一部分是注意力汇聚点(通常保留最初的 4 个 token),另一部分是保留最近 token 的滚动 KV 缓存(Rolling KV Cache)。在注意力计算时,汇聚点 token 吸收了大量无用的注意力权重,其机制可由 SoftMax 的分布特性表示:
$$ \mathrm{SoftMax}(x)_i = \frac{e^{x_i}}{e^{x_1} + \sum_{j=2}^N e^{x_j}} $$其中 $x_1$ 代表初始 token 的注意力对数,它通常远大于其他位置的值($x_1 \gg x_j$),从而稳定了整体分母。为了让这种拼接缓存生效,StreamingLLM 必须绕开绝对位置编码的限制。当为 token 分配相对距离和位置信息时,StreamingLLM 严格根据 token 在当前缓存中的相对索引来分配位置,而不是使用它们在原始长文本中的绝对位置。例如,如果当前缓存中包含原始位置为 [0, 1, 2, 3, 6, 7, 8] 的 token,StreamingLLM 会为它们分配 [0, 1, 2, 3, 4, 5, 6] 的位置编码。对于 RoPE 位置编码,方法是在引入旋转变换之前缓存键(Keys),然后在每次解码阶段对滚动缓存中的键应用位置变换;对于 ALiBi 编码,则是直接对注意力分数应用连续的线性偏置。这种缓存位置对齐机制确保了模型永远不会处理超过其预训练窗口大小的位置索引。此外,作者发现现有模型通常需要 4 个初始 token 才能稳定,因为它们在预训练时没有一个统一的起始 token。为此,作者提出在预训练阶段为所有训练样本的开头添加一个专门的、可学习的“汇聚点 Token”(Sink Token)作为占位符。这个部件的职责是专门吸收冗余的注意力分数,使得未来的模型在流式部署时只需要在缓存中保留这 1 个汇聚点 token 即可维持稳定。
效果如何 实验在单张 NVIDIA A6000 显卡上进行,测试了 Llama-2(最高 70B)、MPT(最高 30B)、Falcon(最高 40B)和 Pythia(最高 12B)等多个主流模型系列。对比基线包括代表全局缓存路线的密集注意力(Dense Attention)、代表朴素淘汰路线的窗口注意力(Window Attention),以及代表高质量但低效路线的带重计算的滑动窗口(Sliding Window with Re-computation)。在 PG19 数据集(包含长篇书籍)的语言建模任务中,密集注意力因显存溢出而失败,窗口注意力在超出缓存时崩溃,而 StreamingLLM 能够在高达 400万 token 的超长文本上保持极其稳定的困惑度,效果几乎与重计算基线完全一致。在效率方面,由于避免了二次方的上下文重计算,StreamingLLM 的单 token 解码速度比重计算基线提升了高达 22.2倍,同时保持了恒定的极低显存占用。在模拟真实场景的 StreamEval 多轮问答测试中,只要问题和答案的距离在缓存窗口内,StreamingLLM 在输入长度接近 12万 token 时依然保持了高准确率。作者也通过从头预训练 160M 参数的模型证实了添加专属 Sink Token 不会损害模型在常规 NLP 基准上的零样本准确率。该方法明确的局限性在于,它仅仅是让模型能够基于缓存内的 token 生成连贯文本,并没有真正扩展模型的上下文窗口或长期记忆。如果任务需要依赖已经被挤出滚动缓存的历史数据(例如长文档摘要或跨度极大的问答),该方法将会失效,因此它主要适用于只需要短期记忆的日常对话等流式场景。
在诸如多轮对话等流式应用中部署大型语言模型(LLMs)有着迫切的需求,因为这些场景通常涉及长时间的交互。然而,这种部署面临两个主要挑战:首先,在解码阶段,缓存先前token的键和值状态(Key and Value states, KV)会消耗大量的内存并增加解码延迟;其次,现有的流行LLMs无法泛化到比其预训练序列长度更长的文本上。虽然仅缓存最近KV的窗口注意力(Window attention)是一种直观的方法,但当文本长度超过缓存大小时,该方法会因为初始token被驱逐而导致模型崩溃。带有重计算的滑动窗口方法虽然性能好,但由于其二次方的计算复杂度而速度极慢。
本文的研究目标是提出一种高效的框架,使经过有限长度注意力窗口训练的LLMs能够在不进行任何微调的情况下,泛化到无限序列长度的文本上。
本文的主要创新点如下:
1. 发现了“注意力汇聚点(Attention Sink)”现象:自回归LLMs会将大量的注意力分数分配给初始的token,即使这些token在语义上并不重要。这是因为Softmax函数要求所有上下文token的注意力分数总和为1,模型倾向于将不必要的注意力值倾倒给全局可见的初始token。
2. 提出了StreamingLLM框架:该框架通过保留注意力汇聚点(即几个初始token)的KV以及滑动窗口的KV,来锚定注意力计算并稳定模型性能。StreamingLLM使Llama-2、MPT、Falcon和Pythia等模型能够稳定高效地处理高达400万甚至更多token的文本。
3. 提出了针对流式部署的预训练策略:发现在预训练期间添加一个占位符token作为专门的注意力汇聚点,可以进一步改善流式部署,使得模型仅需一个专用的汇聚token即可维持性能,而不需要保留多个初始token。
4. 显著的效率提升:在流式设置中,StreamingLLM相比于带有重计算的滑动窗口基线方法,实现了高达$22.2\times$的加速。
窗口注意力的失效起点:图3展示了在20K token文本上的语言建模困惑度(Perplexity)。明显可以看出,当文本长度超越缓存大小,导致初始token被排除时,困惑度会急剧上升。这表明,无论初始token距离正在预测的token有多远,它们对于维持LLMs的稳定性都至关重要。
移除初始KV导致模型崩溃的原因:为了理解原因,作者可视化了Llama-2-7B模型所有层和头的注意力图,如图2所示。作者发现,在底部的两层之上,模型在所有层和头中始终将注意力集中在初始token上。其含义很明确:移除这些初始token的KV将消除注意力计算中SoftMax函数($\mathrm { S o f t M a x } ( x ) _ { i } = \frac { e ^ { x _ { i } } } { e ^ { x _ { 1 } } + \sum _ { j = 2 } ^ { N } e ^ { x _ { j } } } , \quad x _ { 1 } \gg x _ { j } , j \in 2 , \ldots , N$)分母的很大一部分。这种改变会导致注意力分数分布发生重大偏移,使其偏离正常推理设置下的预期。
初始Token重要性的两种解释与验证:关于初始token在语言建模中的重要性有两种可能的解释:(1) 它们的语义至关重要,或者 (2) 模型学习到了对其绝对位置的偏置。为了区分这两种可能性,作者进行了实验,将前四个token替换为换行符 \n。观察结果表明,模型仍然显著强调这些初始的换行符。此外,重新引入它们可将语言建模困惑度恢复到与拥有原始初始token相当的水平。这表明起始token的绝对位置,而不是它们的语义价值,具有更大的重要性。
LLM将初始Token作为注意力汇聚点(Attention Sinks):为了解释为什么模型不成比例地关注初始token(无论它们与语言建模的语义相关性如何),作者引入了“注意力汇聚点(attention sink)”的概念。SoftMax函数的性质防止所有被关注的token具有零值。这就要求在所有层的所有头中聚合来自其他token的一些信息,即使当前的嵌入(embedding)具有足够的自包含信息来进行预测。因此,模型倾向于将不必要的注意力值倾倒给特定的token。在量化异常值领域也观察到了类似的现象(【78,SmoothQuant: Accurate and efficient post-training quantization for large language models+2023+ICML】等),从而导致了SoftMax-Off-by-One(【48,Attention is off by one+2023+URL】)作为潜在补救措施的提出。
为何自回归LLM一致关注初始Token:为什么各种自回归LLM(如Llama-2、MPT、Falcon和Pythia)一致将初始token作为其注意力汇聚点,而不是其他token?作者的解释很直接:由于自回归语言建模的顺序性质,初始token对所有后续token都是可见的,而后面的token只对有限的后续token可见。因此,初始token更容易被训练成注意力汇聚点,捕获不必要的注意力。
多个初始Token作为汇聚点的原因:作者指出LLMs通常被训练为利用多个初始token作为注意力汇聚点,而不是仅仅一个。如图2所示,引入四个初始token作为注意力汇聚点足以恢复LLM的性能。相反,仅仅添加一个或两个并不能实现完全恢复。作者认为出现这种模式是因为这些模型在预训练期间没有在所有输入样本中包含一致的起始token。虽然Llama-2确实在每个段落前加了“<s>”token,但它是在文本分块之前应用的,导致在第零个位置大部分是随机token。这种缺乏统一起始token的情况导致模型使用几个初始token作为注意力汇聚点。作者假设,通过在所有训练样本的开头加入一个稳定的可学习token,它可以单独作为一个专门的注意力汇聚点,消除需要多个初始token来确保一致流式处理的需求。
带有注意力汇聚点的滚动KV缓存(Rolling KV Cache with Attention Sinks):为了在已训练的LLMs中实现流式处理,作者提出了一种简单的方法,可以在不进行任何模型微调的情况下恢复窗口注意力的困惑度。除了当前的滑动窗口token外,作者在注意力计算中重新引入了几个起始token的KV。StreamingLLM中的KV缓存可以在概念上分为两部分,如图4所示:1)注意力汇聚点(四个初始token)用于稳定注意力计算;2)滚动KV缓存保留最近的token,这对语言建模至关重要。StreamingLLM的设计具有通用性,可以无缝集成到任何采用相对位置编码(如RoPE(【66,Roformer: Enhanced transformer with rotary position embedding+2021+arXiv】)和ALiBi(【55,Train short, test long: Attention with linear biases enables input length extrapolation+2022+ICLR】))的自回归语言模型中。
缓存内位置分配机制:在确定相对距离并向token添加位置信息时,StreamingLLM关注的是缓存内的位置,而不是原始文本中的位置。这种区别对StreamingLLM的性能至关重要。例如,如果当前缓存(图4)包含token [0, 1, 2, 3, 6, 7, 8] 并且正在解码第9个token,则分配的位置是 [0, 1, 2, 3, 4, 5, 6, 7],而不是原始文本中的位置,原始文本中的位置将是 [0, 1, 2, 3, 6, 7, 8, 9]。
不同位置编码的集成方式:对于像RoPE这样的编码,作者在引入旋转变换之前缓存token的Keys。然后,在每个解码阶段,作者对滚动缓存中的keys应用位置变换。另一方面,与ALiBi的集成更直接。在这里,连续的线性偏置被应用,而不是对注意力分数应用“跳跃”偏置。这种在缓存内分配位置嵌入的方法对StreamingLLM的功能至关重要,确保模型即使超出其预训练注意力窗口大小也能高效运行。
使用注意力汇聚点预训练LLM(Pre-training LLMs with Attention Sinks):正如第3.1节所阐述的,模型过度关注多个初始token的一个重要原因是缺乏一个指定的汇聚token来卸载过量的注意力分数。由于这个原因,模型无意中使用了全局可见的token(主要是初始token)作为注意力汇聚点。一个潜在的补救措施是故意包含一个全局可训练的注意力汇聚点token,表示为“Sink Token”,它将充当不必要注意力分数的存储库。或者,用类似SoftMax-off-by-One(【48,Attention is off by one+2023+URL】)的变体替换传统的SoftMax函数:$\mathrm { S o f t M a x } _ { 1 } ( x ) _ { i } = \frac { e ^ { x _ { i } } } { 1 + \sum _ { j = 1 } ^ { N } e ^ { x _ { j } } }$,该变体不需要所有上下文token上的注意力分数总和为1,也可能有效。需要注意的是,$\mathbf { S o f t M a x } _ { 1 }$ 等效于在注意力计算中前置一个具有全零Key和Value特征的token。作者将此方法表示为“Zero Sink”以适应其框架。
预训练策略的验证:为了验证上述假设,作者在相同的设置下从头开始预训练了三个具有1.6亿参数的语言模型。第一个模型使用标准的SoftMax注意力(Vanilla),第二个模型用 $\mathbf { S o f t M a x } _ { 1 }$(Zero Sink)替换了常规注意力机制,第三个模型在所有训练样本中前置了一个可学习的占位符token(Sink Token)。结果表明,虽然zero sink在一定程度上缓解了注意力汇聚点问题,但模型仍然依赖其他初始token作为注意力汇聚点。引入sink token在稳定注意力机制方面非常有效。只需将此sink token与最近的token配对,就足以锚定模型的性能,并且最终的评估困惑度甚至略有改善。基于这些发现,作者建议在所有样本中使用sink token来训练未来的LLMs,以优化流式部署。
数据集:
模型与架构关键参数:
硬件配置:
软件配置:
消融实验:
效率结果:图10在Llama-2-7B/13B上基准测试了StreamingLLM的解码延迟和内存使用情况。随着缓存大小的增加,StreamingLLM的解码速度呈线性增长,而重计算基线呈二次上升。StreamingLLM实现了高达$22.2\times$的每token加速,同时保持了与重计算基线一致的内存占用。
在流式应用中部署LLMs有着迫切需求,但由于效率限制和较长文本下性能降低而面临挑战。窗口注意力提供了一种部分解决方案,但当排除初始token时,其性能会暴跌。认识到这些初始token作为“注意力汇聚点(attention sinks)”的作用,本文引入了StreamingLLM——一个简单而高效的框架,使LLMs能够在不进行微调的情况下处理无限长度的文本。通过将注意力汇聚点与最近的token结合起来,StreamingLLM可以高效地建模高达400万token的文本。作者进一步证明,使用专用的sink token预训练模型可以改善流式性能。StreamingLLM首次解耦了LLM的预训练窗口大小与其真实的文本生成长度,为LLMs的流式部署铺平了道路。
相关工作分类与分析:
* 长度外推(Length Extrapolation):旨在使基于短文本训练的语言模型在测试时处理更长文本。主要研究方向是开发相对位置编码(如RoPE(【66,Roformer: Enhanced transformer with rotary position embedding+2021+arXiv】))。尽管ALiBi(【55,Train short, test long: Attention with linear biases enables input length extrapolation+2022+ICLR】)通过基于距离偏置注意力分数展现出更好的外推性,但当文本远超训练长度时仍会崩溃。现有方法尚未实现无限长度外推。StreamingLLM属于这一类,致力于处理无限长度输入,但不扩展注意力窗口。
* 上下文窗口扩展(Context Window Extension):致力于扩大LLM单次前向传递能处理的token数。解决方案包括系统优化(如FlashAttention(【17,FlashAttention: Fast and memory-efficient exact attention with IO-awareness+2022+arXiv】))和近似注意力方法。近期有工作通过位置插值和微调扩展RoPE上下文(【12,Extending context window of large language models via positional interpolation+2023+arXiv】等)。这些技术只在有限程度上扩展窗口,与StreamingLLM处理无限输入的核心目标不同,但两者是正交且可结合的。
* 改善LLM对长文本的利用(Improving LLMs’ Utilization of Long Text):优化LLM以更好地捕获和使用上下文内容。如Liu等人指出,扩大窗口不等于能胜任长上下文利用。StreamingLLM集中于稳定利用最近的token,实现无缝流式应用。
应用与局限性(Appendix A):StreamingLLM特别适合多轮对话等流式应用,这些应用连续运行且不严重依赖大量内存或历史数据。然而,它不扩展模型的上下文窗口或增强长期记忆。模型被限制在当前缓存的范围内运行,因此不适合需要长期记忆和广泛数据依赖的任务(如长文档问答和摘要)。
更广泛的社会影响(Appendix A):StreamingLLM显著提高了LLMs的效率和可访问性,使对话代理中的交互更加无缝,降低了计算负载,符合环保AI的需求。其潜在的负面影响与一般语言模型相同,如错误信息和偏见内容生成的风险。
其他相关工作(Appendix B):
* 稀疏Transformer:如Sparse Transformer(【14,Generating long sequences with sparse transformers+2019】)、LongFormer(【5,Longformer: The long-document transformer+2020+arXiv】)等通过局部窗口或全局token降低复杂度。但它们需要自定义GPU内核,全局注意力不适合自回归模型,且不兼容现有预训练模型。StreamingLLM则易于使用标准GPU内核实现,并兼容预训练自回归模型。
* 同期工作:Han等人提出了“$\Lambda$”形注意力模式来增强长度泛化;Darcet等人在Vision Transformers中观察到了类似的注意力集中在随机背景补丁上的现象,称为“registers”(【18,Vision transformers need registers+2023】)。但StreamingLLM发现的“attention sink”存在于自回归模型的初始token中,表明SoftMax函数在其中起着更基础的作用。
StreamEval中查询-答案距离增加时的准确率(Appendix C):表7评估了Llama-2-7B-32K-Instruct模型在StreamEval上不同查询-答案行距离下的准确率。结果表明,当查询和答案之间的token距离在缓存大小内时,StreamingLLM保持准确率。然而,随着距离的增加,准确率下降,并在超过缓存容量时降至零。这表明StreamingLLM无法扩展上下文长度,也强调了当前模型无法充分利用缓存内上下文信息的挑战。
长距离基准评估(Appendix D):在LongBench(【4,Longbench: A bilingual, multitask benchmark for long context understanding+2023+arXiv】)上的评估(表8)表明,StreamingLLM使用 $4 + 3496$ 的缓存配置表现不如截断基线(保留前后各1750个token),这可能是由于丢失了关键的初始输入提示信息。然而,将注意力汇聚点数量对齐到1750后,性能恢复到文本截断基线的水平,证实了StreamingLLM的有效性取决于其缓存内的信息。
长序列上的注意力可视化定量分析(Appendix E & F):图11可视化了Llama-2-7B在长度为128的较长序列上的注意力,发现初始token的注意力分数远高于其余token。图12进一步定量分析了长度为4096的极长输入。数据代表第4096个token分配给每一层中初始token的注意力。结果显示,除了底部两层外,第一个token的注意力分数非常高,通常超过总注意力的一半,从经验上证实了注意力汇聚点现象在长序列中的存在。
Llama-2-70B注意力可视化(Appendix G):图13可视化了Llama-2-70B的注意力,发现在Llama-2-7B上的观察结果同样适用于70B模型,即在大多数层中,初始token的注意力分数远高于其余token。
编码器Transformer中的注意力汇聚点(Appendix H):作者假设注意力汇聚点现象延伸到BERT(【20,Bert: Pre-training of deep bidirectional transformers for language understanding+2019+NAACL】)和ViT(【21,An image is worth 16x16 words: Transformers for image recognition at scale+2021】)等编码器模型。图14中对BERT-base-uncased的分析表明,该模型在大多数层中对无处不在的[SEP] token分配了不成比例的高注意力分数,证明了注意力汇聚点现象是所有Transformer架构的普遍特征。
预训练阶段使用更多Sink Tokens(Appendix I):图15表明,在预训练期间包含一个或两个sink token的预训练损失曲线与基线模型非常相似。然而,表9和表10详细说明了引入第二个sink token并未在大多数基准任务上产生实质性的性能改善,也没有增强流式性能。模型似乎依赖这两个sink token来维持稳定的流式性能。这表明单个sink token已足够,与ViT中发现多个“registers”有益的情况形成对比。