单卡A100实现百万token推理,速度快10倍,这是微软官方的大模型推理加速
先来看一组数据。微软最新的这项研究,让开发者可以在一台单卡机器上,以10倍的速度处理超过100万token的输入文本。这意味着什么?意味着长上下文AI终于不再是实验品,而是触手可及的生产力工具。
大型语言模型已经全面进入长上下文处理的时代。上下文窗口从之前的128K猛增到10M级别,进步不可谓不大。然而麻烦也随之而来——注意力机制那二次复杂度可不是闹着玩的。模型处理输入提示并开始产生第一个token,也就是我们常说的预填充阶段,动辄就要花上几分钟。首Token生成时间过长,用户体验自然会大打折扣,这也成为长上下文LLM大规模落地的关键瓶颈。
举个例子就清楚了(如图2a所示)。在一台装了A100的机器上跑LLaMA-3-8B模型,如果输入提示有30万个token,预填充阶段需要整整6分钟;等到100万个token时,这一数字直接飙升到30分钟。问题是,自注意力计算的开销占了总预填充延迟的90%以上,这才是LLM处理长上下文的核心瓶颈所在,说它是拦路虎一点不为过。
现有的加速预填充方法,一旦面对长上下文场景,要么精度掉得厉害,要么效率提不上去。为了解决这个顽疾,来自微软和萨里大学的研究团队提出了一种专门面向长序列预填充的稀疏计算方法——
MInference(Million tokens Inference)
- 论文地址:https://arxiv.org/pdf/2407.02490
- 论文主页:https://hqjiang.com/minference.html
- 论文标题:MInference 1.0: Accelerating Pre-filling for Long-Context LLMs via Dynamic Sparse Attention
MInference可以直接用在现有的大模型上,不需要动预训练设置,也不需要额外微调。研究团队在多个下游任务和模型中做了验证,包括InfiniteBench、RULER、PG-19和Needle In A Haystack等评估基准,涉及LLaMA-3-1M、Yi-200K、GLM-4-1M、Phi-3-128K、Qwen2-128K等模型。实验结果是实打实的:MInference将A100上的预填充推理延迟降低了最高10倍,同时保持准确性不掉。
使用MInference 1.0,长上下文LLM(比如LLaMA-3-8B-1M、GLM-4-1M)在单张A100上的推理速度提升了10倍,而且准确度反而还有提升。方法介绍
作者把方法命名为MInference,野心也很直白:希望在一台A100机器上实现百万token级别的推理。这是一种无需训练的高效方法,核心思路是基于动态稀疏注意力来加速预填充阶段。
研究者的关键洞察是:注意力机制,尤其是在长上下文中,是稀疏且动态的。不同输入之间,稀疏模式的差异非常明显。而这种动态稀疏性,最终呈现出三种适用于所有输入的独特空间聚合模式:A形(A-shape)、垂直-斜线(Vertical-Slash)和块状-稀疏(Block-Sparse)。
执行流程其实很清晰。首先用内核感知的稀疏模式搜索算法,为每个注意力头离线确定最佳动态稀疏模式(参见算法1)。推理过程中,根据头部的模式动态逼近动态稀疏指数(算法2、3)。最后,用优化后的GPU内核完成高效的动态稀疏注意力计算,大幅缩短预填充阶段的延迟。
拿"垂直-斜线"模式来说。作者首先利用最后一个Q和K之间的注意力计算,估算出垂直线和斜线的最佳指数,然后借助动态稀疏编译器PIT和Triton构建垂直-斜线FlashAttention内核来加速计算。对于A形、垂直-斜线和块状-稀疏模式,则先在注意力计算中使用Q和K的均值池化,利用均值池化和MatMul的可交换属性估计出块状-稀疏指数,再用Triton构建块稀疏FlashAttention内核。具体内核实现,可以参看附录C.4和代码。
在长上下文基准中的评估结果
研究团队在多种场景中对MInference进行了测试,覆盖了问答、编码、基于检索的任务、多跳问答、总结和数学任务等。RULER基准包含几个复杂的多跳或多针任务,能准确反映模型实际上下文窗口的大小。如表1所示,MInference不仅保留了模型的实际处理能力,还把实际上下文窗口略微扩展到了32K。
再来看InfiniteBench,平均token长度214K,任务分布更广泛。从表2的结果可以清晰看出,与SoTA基线相比,MInference在所有任务上的表现都始终稳定。值得玩味的是,在面对KV检索这类具有挑战性的检索任务时,所有基线模型准确率都跌到了1.2%以下,基本可以说是"靠蒙"。但MInference却成功保留了处理动态KV对检索的能力。
为了进一步验证不同上下文长度、以及关键信息在提示中不同位置的性能表现,研究团队使用"大海捞针"任务做了全面测试。图1的结果很直观:MInference在不同模型、不同上下文窗口、不同提示信息位置下都表现得游刃有余,性能与原始模型相比基本持平甚至略有提升。在LLaMA-3-8B和GLM-4-9B-1M上,MInference在高达1M的上下文中实现了全绿表现。相比之下,StreamingLLM和InfLLM在70K上下文窗口时,提示中间段的性能就跌到了20%以下。
研究团队还用PG-19在语言模型任务中做了测试,token数量达100k。图2的结果清晰显示,MInference保持了LLaMA-3-8B和Yi-9B-200K的困惑度,而所有基线方法都出现了不同程度的下降。注意,使用膨胀和步长配置的StreamingLLM比标准版表现要好一些,但和MInference相比还是有差距。
延迟和内核中的稀疏模式
图3展示了三种注意力模式和FlashAttention的微基准测试结果。Vertical-Slash是三种模式中最慢的,但在1M上下文窗口中,相比FlashAttention依然实现了13倍的加速——这个数字初看可能觉得"才13倍",但在超长上下文场景下,已经是很惊艳的突破了。
图4展示了Vertical-Slash头部内核中的稀疏索引:垂直线通过PIT FlashAttention用1x64块计算,斜线则通过块级FlashAttention用64x64块计算。
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名