Skip to content

大模型的参数量及其计算访存开销的理论分析

推理服务系统的根本目标在于降低时延和提高吞吐量,LLM 推理的优化也是如此。首字时延(Time To First Token, TTFT)和吐字时延(Time Per Output Token, TPOT)就是两个非常重要的指标。如何优化 LLM 推理的这两个指标成为近年来学术界热议的问题。在研究这个问题之前,有必要深入理解 LLM 架构,分析其参数量和计算访存开销。


Transformer[^1] 是一种 Encoder-Decoder 结构的模型,它被认为是 LLM 的基础模型。此后,诸如 Encoder only 模型(BERT[^2])、Decoder only 模型(GPT[^3])、Encoder-Decoder 模型都是在 Transformer 基础之上的变体。但不论是何种结构的 LLM,其内部都主要包含如下 block:

  • Multi-Head Self-Attention:多头自注意力
  • Feed-Forward:前馈网络
  • Add & LayerNorm:层归一化

此外,还包括如下一些 block:

  • Token Embedding:词嵌入
  • Position Embedding:位置嵌入
  • Linear:输出的线性层
  • Softmax:激活函数

LLM 超参数定义如下:

超参数符号表示Transformer 中的大小
dmodelM512
dk(dv)D64
HeadH8
dffF2048

根据论文中的定义,其中 M=HDF=4M

参数量理论计算

MHA

多头自注意力模块的每个 head 都需要 2 次 Self-Attention 的计算,即

Attention(Q,K,V)=Softmax(QKdk)V

在那之前,需要将每个 token 从 dmodel 维映射到 dk(dv) 维。(在 Transformer 模型中,dmodel=512,而 dk=dv=dmodelh=64,这里 h=8。)所以,输入的 Q,K,V 向量都需要进行投影操作,共需 3 个线性变换矩阵 WK,WQ,WV,维度均为 (M, D)

最终将每个 head 拼接起来,得到 dv×h=dmodel 维。这个结果再投影到 dmodel 维,需 1 个线性变换矩阵 WO,维度为 (M, M)。即

MultiHead(Q,K,V)=Concat(head1,...,headh)WOwhere headi=Attention(QWiQ,KWiK,VWiV)

因此,多头自注意力模块的总参数量为

(dmodel×dqkv+dqkv)×3×h+(dmodel×dmodel+dmodel)= 4(dmodel2+dmodel)

dff 很大的情况下,可以近似认为参数量为 4dmodel2

FFN

前馈网络模块由 2 个全连接层,即

FFN(x)=ReLU(xW1+b1)W2+b2

其中隐藏层 dff 的维度一般是 dmodel 的 4 倍。所以参数矩阵 Win,Wout 的维度分别为 (M, F)(F, M)。因此,前馈网络模块的总参数量为

(dmodel×dff+dff)+(dff×dmodel+dmodel)= 2dmodeldff+dmodel+dff

Add & LayerNorm

这一模块包含两个操作,一是 Add,即残差连接;二是 Layer Normalization,即

y=xμσ+ϵ×γ+β

其中,μ 为均值,σ 为标准差,ϵ 为定值,γ 为缩放因子,β 为偏移因子。γ,β 是需要学习的参数,是 (1, M) 维的向量。因此,这部分参数量为

2×dmodel

Embedding

Token Embedding 过程包含参数,参数量取决于训练数据 token 的数量 vocab_size。参数矩阵的维度为 (vocab_size, M)

而 Positional Embedding 采用三角函数计算,即

PE(pos,2i)=sin(pos/100002i/dmodel)PE(pos,2i+1)=cos(pos/100002i/dmodel)

此过程不需要学习,因此没有参数。

Linear & Softmax

从最后一个编码器输出的结果需要经过一个线性层,通过 Softmax 化为概率分布后,选择最大概率的 token 输出。因此,线性层的参数矩阵具有 (M, vocab_size) 维。这部分的参数量为

dmodel×vocabsize

参数量估算

LLM 总参数量为上述各部分参数量之和。不考虑输入和输出,按照 GPT 的架构,即 Decoder only,每个 Decoder 共 1 个 MHA、1 个 FFN 和 2 个 Add & LayerNorm。这样我们可以得到参数量为

4(dmodel2+dmodel)+2dmodeldff+dmodel+dff+2×dmodel×2

这里我们继续简化,令 dff=4dmodel=4d,得到

4(d2+d)+2d×4d+d+4d+2d×2= 12d2+13d

dmodel 较大时,Encoder 和 Decoder 部分的参数量和其他部分相比远远更大。因此,要快速估计 LLM 的参数量,一般可以采用 12dmodel2 作为估计值。

验证参数量

Transformer

按照上述计算方法,我们首先验证 Transformer 的参数量。

首先,计算一个 Encoder 的参数量。一个 Encoder 包含 1 个 MHA block、1 个 FFN block 和 2 个 Add & LayerNorm,即

4(dmodel2+dmodel)+2dmodeldff+dmodel+dff+2×dmodel×2= 4×512×(512+1)+2×512×2048+512+2048+2×512×2= 3,152,384

接下来,计算一个 Decoder 的参数量。一个 Decoder 包含 2 个 MHA block、1 个 FFN block 和 3 个 Add & LayerNorm,即

4(dmodel2+dmodel)×2+2dmodeldff+dmodel+dff+2×dmodel×3= 4×512×(512+1)×2+2×512×2048+512+2048+2×512×3= 4,204,032

上面计算的只是一个 Encoder 或 Decoder 的参数量。Transformer 的 Encoder-Decoder 的数量均为 6 个。

除此之外,从最后一个编码器输出的结果需要经过一个线性层,通过 Softmax 化为概率分布后,选择最大概率的 token 输出。因此,线性层的参数矩阵具有 (M, vocab_size) 维。在 Transformer 论文中,训练所用的英语和德语对照表的 token 数量为 37,000 个。

因此,不考虑 Embedding,Transformer 的参数量为

(3,152,384+4,204,032)×6+37,000×512= 63,082,496

这和论文中给出的 65M 的参数量在同一个量级上。只不过,这里没有考虑输入和输出部分的参数量。如果模型足够大,那么参数量主要取决于 Encoder 和 Decoder 部分。

LLaMA

LLaMA[^4] 是基于 Transformer 的 Decoder only 模型。考虑到 LLaMA 是参数量远大于 Transformer 的 LLM,所以我们采用估算法,可以得到下表:

ModeldmodelLayer采用 12dmodel2 估算采用 12dmodel2+13dmodel 估算
LLaMA-7B4096326,442,450,9446,444,154,880
LLaMA-13B51204012,582,912,00012,585,574,400
LLaMA-30B66566031,897,681,92031,902,873,600
LLaMA-65B81928064,424,509,44064,433,029,120

可见,估算值和论文给定的参数量相差无几。同时也说明,LLM 参数量越大,Encoder 和 Decoder 部分的参数量占比也就越大。可以基本上忽略输入输出部分的参数。

计算访存开销

首先需要给出如下两个定义:

  • FLOPS (Floating Point Operations Per Second): 每秒浮点运算次数,用于衡量模型计算开销。
  • MOPS (Memory Operations Per Second): 每秒内存访问次数,用于衡量模型访存(I/O)开销。

一般地,如果一个系统在单位时间内访存次数越小而计算次数越多,那么该系统的吞吐量就越大。定义算术强度(Arithmetic Intensity[^5])为

Arithmetic Intensity=FLOPSMOPS

显然,算术强度越大,系统的吞吐量越高。

全量推理

全量推理阶段,即输入一个完整的 sequence,到输出第一个 token 的过程。其计算和访存开销与 sequence length 和 batch size 直接相关,分别用 Sin(下面简写为 S)和 B 来分别表示其大小。

对于 MHA 的 4 个 proj 算子,即 WQ,WK,WV,WO,其映射的维度都是一致的,即从 dmodelhead×dk(或反过来)。这部分的 FLOPS 和 B,S,M,H,D 成正比,即

FLOPSBSMHD=BSM2

由于需要把模型权重和中间激活 tensor 都加载到显存中,其 MOPS 由 2 个部分组成,即

MOPS2BSM+MHD=2BSM+M2

其中,常数 2 表示一对显存读写操作,因为中间激活 tensor 需要先读取,再将结果写入显存。

对于 Self-Attention 计算,Q×K 和 Score(用 P 表示)×V 这 2 个线性算子的 FLOPS 与 B,S2,H,D 成正比,即

FLOPSBS2HD=BS2M

不过,由于 Self-Attention 没有模型参数,其 MOPS 由 2 个部分组成,即读取 QK 向量,将结果写入显存,即

MOPS2BSHD+BS2H=2BSM+BS2H

由于 Self-Attention 本质上是计算余弦距离,每个 head 内的向量计算的实际上是内积。这个过程需要额外的多次访存,以存储中间结果。所以,这部分的访存开销和 B,S2,H 成正比。

最后,对于 FFN 的 2 个线性算子,即 Win,Wout 2 个 MLP 层,其 FLOPS 与 B,S,M,F 成正比,即

FLOPSBSMF+BSMF=8BSM2

其中,加号两边的部分虽然一样,但它们分别表示权重矩阵 Win,Wout 和偏置向量 bin,bout

同理,和 MHA 的 4 个 proj 算子类似,FFN 的 MOPS 也由 2 个部分组成,即

MOPSBSM+MF=BSM+4M2

综上所述,对于 Encoder 的 8 个线性算子,其 FLOPS、MOPS 和算术强度如下表[^6]所示:

StageFLOPSMOPSArithmetic Intensity
Q×WQO(BSM2)O(2BSM+M2)O(12M+1BS)
K×WKO(BSM2)O(2BSM+M2)O(12M+1BS)
V×WVO(BSM2)O(2BSM+M2)O(12M+1BS)
Q×KO(BS2M)O(2BSM+BS2H)O(11D+1S)
P×VO(BS2M)O(2BSM+BS2H)O(11D+1S)
A×WOO(BSM2)O(2BSM+M2)O(12M+1BS)
F×WinO(8BSM2)O(BSM+4M2)O(81M+4BS)
F×WoutO(8BSM2)O(BSM+4M2)O(81M+4BS)

上表给出了 LLM 主要的 8 个线性算子的 FLOPS、MOPS 和算术强度,这些开销在 LLM 推理过程中占主导地位。具体来说,这 8 个线性算子可以分为如下两类:

  • proj:即 Projection 投影,是激活矩阵和权重矩阵相乘的过程。包括 MHA block 的 Q×WQ,K×WK,V×WVA×WO 以及 FFN block 的两个 MLP 层。
  • act-to-act:即 MHA block 的两次自注意力计算 Q×KP×V

通过上表,我们可以得出如下几点结论:

  1. 序列长度和批量大小对计算和访存开销的影响成正比。
  2. dmodelproj 算子的 FLOPS 和 MOPS 都具有二次的影响,
  3. 序列长度对 act-to-act 算子的 FLOPS 和 MOPS 都具有二次的影响。在序列长度较小时,act-to-act 算子的计算量较小;但在序列长度较大时,act-to-act 算子的计算量较大。
  4. 增大序列长度、批量大小、dmodel,以及减少 head 的数量,都有助于提高算术强度。不过,序列长度和批量大小受显存限制不可能无限增加,dmodel 和 head 是超参数,可以看作是定值。所以,吞吐量提升的瓶颈在于显存。

增量推理

增量推理阶段,即从输出第一个 token 的过程,自回归地推理出后续的 token,直至最后一个 token 的过程。其计算和访存开销与输出的序列长度直接相关,用 Sout 来表示其大小。对于增量推理阶段,情况比全量推理阶段要更加复杂,我们首先给出结论——对于 Decoder 的 8 个线性算子,其 FLOPS、MOPS 和算术强度如下表所示:

StageFLOPSMOPSArithmetic Intensity
Q×WQi=1SoutO(BM2)i=1SoutO(2BM+M2)O(12M+1B)
K×WKi=1SoutO(BM2)i=1SoutO(2B[Sin+i1]M+M2)Souti=1SoutO(2(Sin+i1)M+1B)
V×WVi=1SoutO(BM2)i=1SoutO(2B[Sin+i1]M+M2)Souti=1SoutO(2(Sin+i1)M+1B)
Q×Ki=1SoutO(B(Sin+i)M)i=1SoutO(BH[(Sin+i)(D+1)+D])Souti=1SoutO(1D+1Sin+i+1)
P×Vi=1SoutO(B(Sin+i)M)i=1SoutO(BH[(Sin+i)(D+1)+D])Souti=1SoutO(1D+1Sin+i+1)
A×WOi=1SoutO(BM2)i=1SoutO(2BM+M2)O(12M+1B)
F×Wini=1SoutO(8BM2)i=1SoutO(BM+4M2)O(81M+4B)
F×Wouti=1SoutO(8BM2)i=1SoutO(BM+4M2)O(81M+4B)

首先注意,增量推理阶段是一个自回归的过程,因此总的 FLOPS、MOPS 和算术强度为每次推理出一个 token 的叠加,这也是表格中求和符号的由来。

其次,这里默认增量推理阶段采用 K/V Cache,因此需要 K/V Cache 的影响。K/V Cache 是一种用空间换时间的策略,因此它减少了 FLOPS,却又增加了 MOPS。具体来说,在第 i 轮迭代时,已经在显存中缓存了长度为 i1+Sin 的 K/V 向量。此时,只需要计算当前 token 在 WK,WV 上的投影,故此时 FLOPS 的计算复杂度与序列长度无关,即

O(BM2)

相应地,在第 i 轮迭代时,需要从显存中读取长度为 i1+Sin 的 K/V 向量和计算当前 token 在 WK,WV 上的投影所需的权重矩阵,即

O(2B[Sin+i1]M+M2)

在后续的 MHA 计算中,第 i 轮迭代的序列长度就变为 Sin+i。计算完 MHA 后,序列长度就恒为 1 了,故增量推理阶段 FFN block 的相关指标与全量推理阶段相同。

通过上表,我们可以得出如下几点新的结论:

  1. 输入序列长度的增加直接导致 K/V Cache 线性增加,于是访存开销随之增加。这样会降低 proj 的算术强度,但是却能够提升 act-to-act 的算术强度。

  2. 输出序列长度的增加直接导致自回归解码迭代次数的增加,对计算和访存的影响一定会增加。除了对 K/V 向量的 projact-to-act 算子的 FLOPS 和 MOPS 具有二次的影响外,对其他算子的 FLOPS 和 MOPS 只是成倍增加。

  3. 和全量推理的情况一样,增加批量大小和 dmodel 和减少 head 的数量,都有助于提高算术强度。

  4. 输出序列长度对算术强度会有怎样的影响?下面分别就 K/V 向量的 projact-to-act 算子的算术强度进行推导。

Souti=1SoutO(Sin+i1M+1B)= Souti=1SoutO(Sin+i1M+1B)= SoutO((Sin1)SoutM+Sout(Sout+1)2M+SoutB)= O(2Sin+Sout12M+1B)

所以,输出序列长度的增加其实可以提高 K/V 向量 proj 的算术强度。

不过,如果要对 act-to-act 算子进行算术强度分析,涉及到如下的数列求和问题:

i=1Sout1i+Sin

这相当于求调和级数的前 Sout 项与前 Sin 项的差。结合调和级数的前 n 项和公式,得到

S(n)=i=1n1ii=1Sout1i+Sin=S(Sout)S(Sin)ln(Sout+1)ln(Sin+1)=lnSout+1Sin+1

所以,act-to-act 算子进行算术强度为

Souti=1SoutO(1D+1Sin+i+1)= SoutO(SoutD+lnSout+1Sin+1+Sout)= O(11D+1+1SoutlnSout+1Sin+1)

容易判断其单调性,输出序列长度的增加同样也会增加 act-to-act 算子的算术强度。

[^1]: Vaswani, Ashish, et al. “Attention Is All You Need.” Proceedings of the 31st International Conference on Neural Information Processing Systems, Curran Associates Inc., 2017, pp. 6000–10. [^2]: Devlin, Jacob, et al. “BERT: Pre-Training of Deep Bidirectional Transformers for Language Understanding.” arXiv:1810.04805 [Cs], May 2019. [^3]: Radford, Alec, and Karthik Narasimhan. Improving Language Understanding by Generative Pre-Training. 2018. [^4]: Touvron, Hugo, et al. LLaMA: Open and Efficient Foundation Language Models. arXiv:2302.13971, arXiv, 27 Feb. 2023. [^5]: Kim, Sehoon, et al. Full Stack Optimization of Transformer Inference: A Survey. arXiv:2302.14017, arXiv, 27 Feb. 2023. [^6]: 本表根据 剖析 GPT 推断中的批处理效应 一文进行整理。下同。

最近更新