首页 > 教程攻略 > ai资讯 >用Deja Vu来提高 LLM 的推理速度

用Deja Vu来提高 LLM 的推理速度

来源:互联网 时间:2026-07-30 13:09:54

ICML23上有一篇很有意思的工作,核心思路是在推理阶段,根据当前的输入X,动态地挑选一部分网络参数参与计算,而不是每次都动用全部参数。从某种角度看,这可以理解成一种“动态剪枝”策略。苹果最近推出的“LLM in flash”方案,差不多就是对这项技术的直接应用。从这能看出,Deja Vu未来在移动端部署大语言模型这件事上,潜力相当可观。

更关键的是,这个方法实现起来并不复杂。

在算力极度受限的场景下,提升大语言模型的推理速度,现实意义非常突出。剪枝,作为模型轻量化的重要途径,其核心前提是模型参数存在稀疏性——也就是说,我们可以裁掉一部分不那么关键的参数,在尽量不影响模型表现的前提下,把参数量降下来。

实际应用中,通常的做法是用特定的剪枝算法,从一个较大的模型(叫它模型A)出发,剪出一个较小的模型(模型B)。推理阶段就用模型B来完成计算,过程如图1所示。

在这种模式下,对于任何输入X,模型B的产生过程都和X本身无关。这类方法可以叫做“静态剪枝”,其背后的稀疏性假设就是“静态稀疏性”。

而Deja Vu的做法不一样,在这里,模型B的产生是与输入X相关的。所以它是一种“动态剪枝”,对应的假设是“动态稀疏性”。原论文里用的术语是Contextual Sparsity,强调的正是模型B的产生与输入X挂钩,如图2所示。

图2清晰地展示了模型B的产生过程受输入X影响。Deja Vu要解决的核心问题就是:

如何基于输入X,从预训练好的大模型A中,快速“变出”一个针对性的小模型B

1. Contextual Sparsity

Contextual Sparsity的核心思想,在于强调对于已经训练好的大模型A,其参数的重要性是依赖于输入X的。举个例子,对于两个不同的输入X1和X2,通过图2所示的方式剪枝得到的模型B是不一样的。这就是“上下文”的含义——模型B依赖于上下文,而这个上下文指的就是模型的输入。

那么一个关键问题是:Contextual Sparsity在大模型里真的存在吗?

论文里验证的方法相当直接。首先,用输入X做一次前向推理,过程中记录下那些输出具有较大L2范数的MHA(多头注意力)中的head,以及MLP中的神经元。

具体实现起来也不复杂。以MHA为例,它的每个head输出都是一个矩阵,只需要对这个矩阵计算L2范数,然后从所有head中挑出范数最大的那几个就行,如下图所示。

在图1.1的示例里,输入X的序列长度N=10,维度d=6。假设MHA的head数量是3,那么MHA的输出就对应3个不同的head,图中用不同颜色区分。只需计算每个head对应输出的L2范数,然后再找出较大的head。

图1.1中的MHA忽略了最后的线性层

对于MLP,情况稍有不同。对于某个特定的token,MLP的输出是一个向量。但这时候不能直接算整个向量的L2范数,而是要看哪个输出维度的L2范数更大。这是因为

整个向量是由所有神经元共同算出来的

,而每个维度正好对应一个神经元,如下图所示。

MLP是逐位置执行的,所以它的物理意义应该用图中红色框的形式来理解:每一列对应输入中每一个token的变换。但在找有效神经元时,需要按行来挑选。换句话说,要计算每一行的L2范数,然后挑出范数较大的神经元。

找到这些L2范数较大的head和神经元后,论文里用同样的输入X再做一次前向推理,但这次只让被挑出来的部分head和MLP神经元参与计算。结果发现,仅用这些经过挑选的组件,几乎不影响模型的效果。

也就是说,通过这个简单的实验,作者们确认了大语言模型中确实存在Contextual Sparsity:模型中有一些与输入强相关的高效参数。只用这些参数,就能达到和全参数模型几乎一致的表现。

再重复一遍,不同的输入X,对应的高效参数部分是不一样的。这和传统的剪枝方法有本质区别,也正是Contextual Sparsity这个名字的由来。

注:Transformer原文中,MHA每个head的输出是通过拼接成一个与输入尺寸一致的张量,再接全连接层做变换。这篇论文因为需要挑选部分head,流程上稍有调整:每个head的输出会先接一个全连接层,变换成与输入尺寸一致的张量;然后对所有head的输出求平均和。这样一来,无论选多少个head,MHA的输出尺寸都是一样的。

前面提到,可以通过计算Attention head和MLP神经元输出的L2范数来找到“高效参数”。那么,这个比例大概是多少呢?比如,我们算出了所有head和所有神经元的L2范数,究竟要选排名前多少的,才能作为“高效参数”?

论文基于OPT模型的实验结果是这样的:

  • Attention head的稀疏率大概在80%左右

  • MLP中神经元的稀疏率大概在95%左右

这意味着,在实际推理时,我们只需要用大约20%的Attention head和大约5%的MLP神经元,就能达到和全参数模型差不多的效果。

2. 稀疏性预测

要利用前面提到的稀疏性来加速推理,就需要有方法能提前、准确地预测,对于当前的输入X,哪些head和哪些MLP神经元是“高效参数”。

这里需要两个预测模型。一个用来预测MHA里哪些head是“高效的”;另一个用来预测MLP里哪些神经元是“高效的”(这里的“神经元”实际上指向参数矩阵的某一列或某一行,具体取决于参数矩阵的定义方式)。

在Deja Vu中,这两个模型的实现都用了一个两层MLP。

以预测Attention head编号为例,假设head数量是256,那么只要把MLP的输出层大小设为256,并为每一个输出加上sigmoid做一个二分类就行(选择或不选择)。

训练数据来自一个完整训练好的大模型。在这个模型推理的过程中,记录下它的Attention输入和Attention输出,算出不同head的L2范数,然后基于一个L2范数的阈值t,把head分成正例和负例。

预测MLP中需要选择的神经元编号,思路和上面基本一致。

下面以一个Transformer模块为例,展示一种朴素实现方法。

公式(1)到公式(4)的流程在逻辑上没问题,但它在原本的流程里额外加了两个步骤:预测head编号(公式1)和预测MLP神经元编号(公式3)。这可能会导致网络的整体速度比原来用全参数时还慢!

下一部分会介绍论文里采用的高效实现方案。

3. 高效实现

基于MLP的稀疏性预测必须做到尽可能高效,否则整个大模型的推理时间可能会因为引入额外的MLP而变得更慢。

这部分整理一下论文里用到的一些高效实现策略。

3.1 并行化稀疏性预测

先说方案,再说原因。

因此原本只能串行执行的四个步骤(公式1到公式4),其中预测编号的两步可以和剩下的两步并行执行。

但为什么可以这样操作?

论文给出的理由是:大语言模型中token的embedding变化非常缓慢。

因为变化慢,所以提前一步去预测似乎也合情合理。论文用两张图说明了token embedding变化缓慢的现象。

图3.1中的左图,展示的是

连续两个网络层之间token embedding的余弦相似度

,高得离谱。右图则是

间隔n层的token embedding之间的余弦相似度

。两张图都很直观地表明,大语言模型里token的Embedding变化确实非常缓慢。因此,用前一层的输入去提前预测下一层的两个编号,是合理的。

3.2 Kernel Fusion

Kernel Fusion算是绝大多数优化工作中的标准操作。在具体动手优化前,首先得想清楚Kernel Fusion是否可行。

在PyTorch里实现论文中的稀疏矩阵乘法,需要先用预测的编号索引,从参数矩阵里取出对应的参数。这会带来3次I/O:1) 读参数矩阵W;2) 取完对应索引后写参数矩阵W1;3) 读W1然后做矩阵乘法。

很明显,对于当前场景,步骤2和步骤3是多余的操作。但在PyTorch提供的现有算子下,只能这样执行。

所以一个很直接的策略就是单独写一个kernel,把步骤1、2、3合并到一起,这样I/O只有一次。

Kernel Fusion这一步,速度提升了4倍(仅指这一步计算的速度)。

3.3 Memory coalescing

论文里给的例子感觉不太友好,和文中的符号有点对不上。

4. 其它

论文里详细的实验结果就不一一列举了。基于OPT模型,在75%的稀疏率下,Deja Vu实现了2到6倍的加速,而且没有掉点,效果相当可观。

唯一的遗憾是,与很多大语言模型领域的论文相比,实验量还是偏小了一点。

这不禁让人联想到神经网络的过参数化问题。真是一个有趣的话题,很想多聊几句,但想多了容易乱,暂且点到为止。