用Deja Vu来提高 LLM 的推理速度
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之间的余弦相似度
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倍的加速,而且没有掉点,效果相当可观。
唯一的遗憾是,与很多大语言模型领域的论文相比,实验量还是偏小了一点。
这不禁让人联想到神经网络的过参数化问题。真是一个有趣的话题,很想多聊几句,但想多了容易乱,暂且点到为止。
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名