《一个Token的旅程》token是如何在大模型中从头到尾走过系列二
提示:文章写完后,目录可以自动生成,如何生成可参考右边的帮助文档
文章目录
前言
提示:这里可以添加本文要记录的大概内容:
基于一的前情提要接着往下说
提示:以下是本篇文章正文内容,下面案例可供参考
一、整体流程以及假设大致介绍
当前商用大模型的推理分为两个步骤,即Prefill和decode。Prefill阶段,用户输入的tokens将全部传入到大模型里面全部一起推理计算,该过程是为了初始化KVcache,以及获得输出的第一个token用来做decode阶段。decode阶段则是基于输出的token做自回归生成。显而易见的是prefill阶段是计算密集型,因为是把用户输入的tokens一次性全部计算完,而decode阶段是显存密集,因为自回归每一个token 的生成都需要和过往全部KV进行计算(其实Prefill也是要做kv计算的,用户输入的tokens中当前token需要跟之前所有token的kv做计算,只是可能因为一次性计算tokens计算量相比显存更加突出,所以才会说prefill是计算密集吧)。所以很多时候我们都能听到PD分离这个概念,不过这里不会对PD分离做过多讲解。
接下来是假设:
用户输入“hello world”。目标长度为5.
一中的核心部分如图:
这部分也是我讲解的重点。
二、按照onnx图走流程
1.input_ids阶段
当前的hello world这两个tokens在输入之前需要根据目标长度做填充从而保证输入的一致性。一般来讲,这是为了用在多batch的情况下才设置目标长度的,当前我们这个背景算是单用户,没有多batch的情况其实是可以不用目标长度这个概念的。目标长度的设定策略一般是这样的:大模型有一个默认的目标长度,当前所有batch中假如有超过该长度的batch,则选择截断该batch,同时确定目标长度为默认目标长度。而假如没有超过默认长度的batch,则会选择最长的batch作为目标长度,剩下的batch不足目标长度的就补充padding。
[hello, world]这边根据前面所述会转换成[bos,hello,world,padding,padding],bos一般是当作起始的意思。
确定好输入后,就需要根据token对应的id去取号了。token和id的对应关系在tokenizer.json和tokenizer_config.json文件里面
在该文件中找到索引如下:
BOS: <|begin_of_text|> (ID: 128000),
“hello”: 15339,
“world”: 14957,
PAD: <|finetune_right_pad_id|> (ID: 128004)
所以[bos,hello,world,padding,padding]转换成[128000,15339,14957,128004,128004]。input_ids阶段结束。
2.Gather阶段
该阶段是将tokens通过embedding矩阵转换为能输入大模型的编码结构,embedding和过往one-hot的方式不一样,embedding在模型训练阶段就已经能感知不同token之间联系的能力,embedding矩阵中每一列都不是单独的个体,举个例子:“king”-“man”+"woman"得到的embedding向量和"queen“的embedding向量会非常相似。
图中可以看到embedding矩阵的大小是[128256,3072],这个是按照[vocab_size,embedding]来的,vocab_size是词元表的大小。前面的[128000,15339,14957,128004,128004]通过该embedding矩阵寻找对应的embedding向量,得到输入为[batch_size,sequence,embedding]为[1,5,3072]。
3.SimplifiedLayerNormalization阶段(RMSNORM)
该阶段是为了做归一化的,归一化的作用是为了保证矩阵中向量里的元素不会偏差过大,举个例子,[120000,0.00012]这样的向量是十分不利于计算的,我们需要把数据经过处理后限制在需要的一定范围内。
RMSNORM的计算公式:
RMSNorm ( x i ) = x i 1 n ∑ j = 1 n x j 2 + ϵ ⋅ γ i \text{RMSNorm}(x_i) = \frac{x_i}{\sqrt{\frac{1}{n}\sum_{j=1}^{n} x_j^2 + \epsilon}} \cdot \gamma_i RMSNorm(xi)=n1∑j=1nxj2+ϵxi⋅γi
而这部分的参数如下图所示:

gamma在这个weight权重里面记录着。 ϵ \epsilon ϵ则示图中的epsilon部分。
注意,RMSNORM计算公式里面计算都是按照一个embedding向量来的,而不是整个输入矩阵。gamma在里面是做点乘。
1 1 n ∑ j = 1 n x j 2 + ϵ \frac{1}{\sqrt{\frac{1}{n}\sum_{j=1}^{n} x_j^2 + \epsilon}} n1∑j=1nxj2+ϵ1
该部分得到为常数。
再乘上对应的 x i x_i xi和 γ i \gamma_i γi也是常数。
所以输入的维度其实页没变化。从该layer输出的数据为[1,5,3072].
4.归一化后的矩阵乘法

可以发现,这里的维度出现了变化,输入为[1,5,3072]经过这个matmul后为[1,5,5120]。这里其实是计算QKV的步骤,[1,5,5120]需要拆开来,得到Q为[1,5,3072],KV各自为[1,5,1024]。为了后面的多头注意力做准备,可以发现Q中3072是KV中1024的三倍,所以后面的GroupQueryAttention(GQA)中是分了三个group。
5.attention_mask

在讲GroupQueryAttention之前,还需要讲一下attention_mask,因为GroupQueryAttention的输入里面还需要attention_mask这边的输入。
attention_mask是在用户输入tokens的刚开始阶段确定的。前面提到输入为[bos,hello,world,padding,padding],这时候attention_mask为[1,1,1,0,0]。这个是用来记录输入中有多少个有意义的token。ReduceSum是记录attention_mask中有几个有效token,由[1,1,1,0,0]得sequence_length为3,sub是用目标长度减去sequence_length,即5-3=2.另一边shape是获取attention_mask矩阵的形状,得到shape[1,5]。然后gather是获得shape矩阵的第二个维度5.所以GQA中从外面传入的两个箭头一个是sub=2,一个是目标长度5(感觉过程有点冗余,不是很懂。)
6.GroupQueryAttention
GQA这边的输入首先要计算ROPE。
1.分组
GQA的参数如下:
do_rotary说明是需要使用ROPE,kv_num_heads=8说明KV分8个头。num_heads说明Q用24个头,scale为注意力计算中的 d k \sqrt{d_k} dk。
说明Q由[1,5,3072]转换到[1,5,24,128],KV由[1,5,1024]转换成[1,5,8,128]。
2.ROPE初始化
ROPE的初始化如下
θ i = base − 2 i d , i = 0 , 1 , 2 , … , d 2 − 1 \theta_i = \text{base}^{-\frac{2i}{d}}, \quad i = 0, 1, 2, \ldots, \frac{d}{2}-1 θi=base−d2i,i=0,1,2,…,2d−1
其中 base = 10000 \text{base} = 10000 base=10000, d = 128 d=128 d=128 是头维度(head dimension), i i i 是维度索引
所以得到的 θ \theta θ矩阵为[1,64]。
下面是需要根据token位置来做乘法
p o s i t i o n = [ 0 , 1 , 2 , … , maxsequence − 1 ] T \mathbf{position} = [0, 1, 2, \ldots, \text{maxsequence}-1]^T position=[0,1,2,…,maxsequence−1]T
这里的position的大小是根据模型自己的默认设置决定的,因为cos和sin的cache都是需要在模型初始化的时候就需要生成好的。
可以看到,模型中初始化cos和sin最长为131072,说明该大模型的输入不能超过131072个token,不然就无法计算。
接下来就是需要构造position和 θ \theta θ的矩阵 M M M:
M = position ⊗ θ = [ 0 , 1 , 2 , … , 131071 ] T ⊗ [ θ 0 , θ 1 , … , θ 63 ] = R 131072 ⊗ R 64 = R 131072 × 64 \begin{aligned} \mathbf{M} &= \text{position} \otimes \boldsymbol{\theta} \\ &= [0, 1, 2, \ldots, 131071]^T \otimes [\theta_0, \theta_1, \ldots, \theta_{63}] \\ &= \mathbb{R}^{131072} \otimes \mathbb{R}^{64} \\ &= \mathbb{R}^{131072 \times 64} \end{aligned} M=position⊗θ=[0,1,2,…,131071]T⊗[θ0,θ1,…,θ63]=R131072⊗R64=R131072×64
接着就可以计算图中GQA中展示的cos_cache和sin_cache了。
c o s _ c a c h e [ p o s i t i o n , i ] = c o s ( M [ p o s i t i o n , i ] ) s i n _ c a c h e [ p o s i t i o n , i ] = s i n ( M [ p o s i t i o n , i ] ) \begin{aligned} cos\_cache[position,i]=cos(M[position,i])\\ sin\_cache[position,i]=sin(M[position,i]) \end{aligned} cos_cache[position,i]=cos(M[position,i])sin_cache[position,i]=sin(M[position,i])
此时初始化结束且cos和sin的规模与onnx上相同
3.ROPE计算
ROPE的计算过程,首先初始化好旋转矩阵
R ( j ) = ( cos ( θ j ) − sin ( θ j ) sin ( θ j ) cos ( θ j ) ) \mathbf{R}^{(j)} = \begin{pmatrix} \cos(\theta_j) & -\sin(\theta_j) \\ \sin(\theta_j) & \cos(\theta_j) \end{pmatrix} R(j)=(cos(θj)sin(θj)−sin(θj)cos(θj))
然后计算公式如下:
[ x , , y , ] = ( cos ( θ j ) − sin ( θ j ) sin ( θ j ) cos ( θ j ) ) × [ x , y ] [x^,,y^,]=\begin{pmatrix} \cos(\theta_j) & -\sin(\theta_j) \\ \sin(\theta_j) & \cos(\theta_j) \end{pmatrix} \times [x,y] [x,,y,]=(cos(θj)sin(θj)−sin(θj)cos(θj))×[x,y]
下面实际演示一下过程。
首先调整KV的格式。KV分组后的格式为[1,5,8,128],为了便于ROPE计算,需要拆一下,变成[1,5,8,64,2]。然后跟cos_cache和sin_cache做乘法(Q也同理):
[ K [ 1 , p o s i t i o n , h e a d , i , 0 ] , K [ 1 , p o s i t i o n , h e a d , i , 1 ] ] = ( cos _ c a c h e [ p o s i t i o n , i ] − sin _ c a c h e [ p o s i t i o n , i ] sin _ c a c h e [ p o s i t i o n , i ] cos _ c a c h e [ p o s i t i o n , i ] ) × [ K [ 1 , p o s i t i o n , h e a d , i , 0 ] , K [ 1 , p o s i t i o n , h e a d , i , 1 ] ] T \begin{aligned} [K[1,position,head,i,0],K[1,position,head,i,1]]=\begin{pmatrix} \cos\_cache[position,i] & -\sin\_cache[position,i] \\ \sin\_cache[position,i] & \cos\_cache[position,i] \end{pmatrix} \times [K[1,position,head,i,0],K[1,position,head,i,1]]^T \end{aligned} [K[1,position,head,i,0],K[1,position,head,i,1]]=(cos_cache[position,i]sin_cache[position,i]−sin_cache[position,i]cos_cache[position,i])×[K[1,position,head,i,0],K[1,position,head,i,1]]T
得到新的K[1,5,8,64,2],再整合成K[1,5,8,128](Q同理).此时ROPE计算完成。
7.构建casual_mask
casual_mask是为了在后面注意力计算的时候对 Q K T QK^T QKT后做掩码遮掩的,这一步是为了遮挡住padding以及当前token往后的token(满足时时序因果,因为不能预知未来知道当前token的下一个token是什么)
前面由attention_mask计算得到的sub=2和attention_mask的size=5就开始起作用了。
我们要构造的casual_mask矩阵的大小是[5,5],其中padding有两个,最终该矩阵如下:
c a s u a l _ m a s k = [ 1 0 0 0 0 1 1 0 0 0 1 1 1 0 0 0 0 0 0 0 0 0 0 0 0 ] casual\_mask=\begin{bmatrix} 1 & 0 & 0 & 0 & 0 \\ 1 & 1 & 0 & 0 & 0 \\ 1 & 1 & 1 & 0 & 0 \\ 0 & 0 & 0 & 0 & 0 \\ 0 & 0 & 0 & 0 & 0 \end{bmatrix} casual_mask=
1110001100001000000000000
然后还需要对该矩阵进行一次转换得到:
c a s u a l _ m a s k = [ 0 − ∞ − ∞ − ∞ − ∞ 0 0 − ∞ − ∞ − ∞ 0 0 0 − ∞ − ∞ − ∞ − ∞ − ∞ − ∞ − ∞ − ∞ − ∞ − ∞ − ∞ − ∞ ] casual\_mask=\begin{bmatrix} 0 & -\infty & -\infty & -\infty & -\infty \\ 0 & 0 & -\infty & -\infty & -\infty \\ 0 & 0 & 0 & -\infty & -\infty \\ -\infty & -\infty & -\infty & -\infty & -\infty \\ -\infty & -\infty & -\infty & -\infty & -\infty \end{bmatrix} casual_mask=
000−∞−∞−∞00−∞−∞−∞−∞0−∞−∞−∞−∞−∞−∞−∞−∞−∞−∞−∞−∞
casual_mask部分结束
8.注意力计算
此时的Q[1,5,24,128], KV[1,5,8,128]。为了便于计算,KV还需要扩展。KV的扩展时按组扩展,前面根据参数分析可得GQA时分成三组了,所以按组扩展的时候KV的每个头需要x3复制才行,最终得到KV[1,5,24,128]。
为了并行计算,seq和head需要互换一下,这里是转置的方式互换,不影响正常计算。得到Q[1,24,5,128],KV[1,24,5,128]。接下来开始注意力计算。注意力计算的公式如下:
Attention ( Q , K , V ) = softmax ( Q K T d k ) V \text{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \text{softmax}\left(\frac{\mathbf{Q} \mathbf{K}^T}{\sqrt{d_k}}\right) \mathbf{V} Attention(Q,K,V)=softmax(dkQKT)V
其中
Q K T = [ 1 , 24 , 5 , 128 ] × [ 1 , 24 , 5 , 128 ] T = [ 1 , 24 , 5 , 5 ] QK^T=[1,24,5,128] \times [1,24,5,128]^T = [1,24,5,5] QKT=[1,24,5,128]×[1,24,5,128]T=[1,24,5,5]
这时候需要将前面的casual_mask引入。
O U T _ m a s k ( h e a d ) [ i , j ] = Q K T ( h e a d ) [ i , j ] + c a s u a l _ m a s k [ i , j ] , i , j = 0 , 1 , 2 , 3 , 4 , 5 OUT\_mask(head)[i,j] = QK^T(head)[i,j] + casual\_mask[i,j], \quad i,j = 0, 1, 2, 3,4,5 OUT_mask(head)[i,j]=QKT(head)[i,j]+casual_mask[i,j],i,j=0,1,2,3,4,5
这样能将每个token对应的未来token位置以及padding位置掩盖。当OUT矩阵进入softmax计算后 − ∞ -\infty −∞部分会变成0,而0部分无影响。这里的 d k \sqrt{d_k} dk就是前面的scale。计算后维度不变。注意,这里的softmax也是按OUT矩阵每行各自softmax,而不是整个OUT矩阵softmax。softmax后维度不变。
此时做最后的计算:
O u t p u t = s o f t m a x o u t × V = [ 1 , 24 , 5 , 5 ] × [ 1 , 24 , 5 , 128 ] = [ 1 , 24 , 5 , 128 ] Output = softmax_{out} \times V=[1,24,5,5] \times [1,24,5,128] = [1,24,5,128] Output=softmaxout×V=[1,24,5,5]×[1,24,5,128]=[1,24,5,128]
然后seq和head转置回来,再多头合并得到Output[1,24,3072]。此刻注意力计算结束。
当前prefill阶段的有效token的KV都会记录在KVcache中。
7.后续计算
后面的计算逻辑都会很简单。
注意力往后做矩阵乘法,维度不变还是[1,5,24,3072]。
下一个是SkipSimplifiedLayerNormalization阶段。该阶段是将未进行注意力计算的输入和注意力计算后的输入一起做个加法(其实就是残差)然后归一化,该做法是为了保证后面的计算不要丢失原本输入的数据特征,要在原本的数据特征的基础上进行拟合。
y = LayerNorm ( x + GQA ( x ) ) \mathbf{y} = \text{LayerNorm}(\mathbf{x} + \text{GQA}(\mathbf{x})) y=LayerNorm(x+GQA(x))
这里的 LayerNorm \text{LayerNorm} LayerNorm根据onnx上的layer表述应该和前面的归一化方法是一样的都是RMSNORM。维度不变还是[1,5,3072]。
后面的计算快速过一下,都是FFN的知识。最终到下一个SkipSimplifiedLayerNormalization的两个输入的维度为[1,5,3072]。
至此,该onnx中基本讲完。onnx图中剩下的很多都是以上环节的重复。
8.末尾环节

最后这里FFN后还需要一次矩阵计算。这里的MatMul是有说法的,128256是该onnx的vocab_size,这里要做的是将embedding转换成token回来。目前这里常用的一个思路是该权重矩阵和开头的gather部分的embedding做权重共享,这是很多流行大模型的常用操作。最终得到输出[1,5,128256]。这时候再做softmax得到每个单词作为下一个token的概率。
但是得到概率并不意味着一定要选择概率最大的作为下一个token输出。这里有不同的解码策略。
1.Greedy Decoding(贪心解码)
logits = h W vocab + b next_token = arg max i logits i = arg max i { logit 1 , logit 2 , … , logit V } \begin{align} \text{logits} &= \mathbf{h} W_{\text{vocab}} + \mathbf{b} \\ \text{next\_token} &= \arg\max_{i} \text{logits}_i \\ &= \arg\max_{i} \{ \text{logit}_1, \text{logit}_2, \ldots, \text{logit}_V \} \end{align} logitsnext_token=hWvocab+b=argimaxlogitsi=argimax{logit1,logit2,…,logitV}
这是选择概率最大的词,但是容易产生重复、单调的文本。
2.Temperature Scaling
scaled_logits i = logits i τ P ( token i ) = exp ( scaled_logits i ) ∑ j = 1 V exp ( scaled_logits j ) = exp ( logits i / τ ) ∑ j = 1 V exp ( logits j / τ ) \begin{align} \text{scaled\_logits}_i &= \frac{\text{logits}_i}{\tau} \\ P(\text{token}_i) &= \frac{\exp(\text{scaled\_logits}_i)}{\sum_{j=1}^{V} \exp(\text{scaled\_logits}_j)} \\ &= \frac{\exp(\text{logits}_i / \tau)}{\sum_{j=1}^{V} \exp(\text{logits}_j / \tau)} \end{align} scaled_logitsiP(tokeni)=τlogitsi=∑j=1Vexp(scaled_logitsj)exp(scaled_logitsi)=∑j=1Vexp(logitsj/τ)exp(logitsi/τ)
temperature < 1: 更确定性(sharper分布)
temperature > 1: 更随机性(flatter分布)
3.Top-k Sampling
V k = { top-k tokens by logits } logits i ′ = { logits i if i ∈ V k − ∞ otherwise P ( token i ) = exp ( logits i ′ ) ∑ j ∈ V k exp ( logits j ′ ) next_token ∼ Multinomial ( P ) \begin{align} \mathcal{V}_k &= \{ \text{top-k tokens by logits} \} \\ \text{logits}'_i &= \begin{cases} \text{logits}_i & \text{if } i \in \mathcal{V}_k \\ -\infty & \text{otherwise} \end{cases} \\ P(\text{token}_i) &= \frac{\exp(\text{logits}'_i)}{\sum_{j \in \mathcal{V}_k} \exp(\text{logits}'_j)} \\ \text{next\_token} &\sim \text{Multinomial}(P) \end{align} Vklogitsi′P(tokeni)next_token={top-k tokens by logits}={logitsi−∞if i∈Vkotherwise=∑j∈Vkexp(logitsj′)exp(logitsi′)∼Multinomial(P)
只考虑概率最高的 k 个词从中按概率随机采样.
4.Top-p (Nucleus) Sampling
sorted_probs = sort ( P ( token i ) , descending = True ) cumsum_probs i = ∑ j = 1 i sorted_probs j V p = { i : cumsum_probs i ≤ p } P ′ ( token i ) = { P ( token i ) ∑ j ∈ V p P ( token j ) if i ∈ V p 0 otherwise next_token ∼ Multinomial ( P ′ ) \begin{align} \text{sorted\_probs} &= \text{sort}(P(\text{token}_i), \text{descending}=\text{True}) \\ \text{cumsum\_probs}_i &= \sum_{j=1}^{i} \text{sorted\_probs}_j \\ \mathcal{V}_p &= \{ i : \text{cumsum\_probs}_i \leq p \} \\ P'(\text{token}_i) &= \begin{cases} \frac{P(\text{token}_i)}{\sum_{j \in \mathcal{V}_p} P(\text{token}_j)} & \text{if } i \in \mathcal{V}_p \\ 0 & \text{otherwise} \end{cases} \\ \text{next\_token} &\sim \text{Multinomial}(P') \end{align} sorted_probscumsum_probsiVpP′(tokeni)next_token=sort(P(tokeni),descending=True)=j=1∑isorted_probsj={i:cumsum_probsi≤p}={∑j∈VpP(tokenj)P(tokeni)0if i∈Vpotherwise∼Multinomial(P′)
选择累积概率达到 p 的最小词集,动态调整候选词数量。
9.自回归环节
前面的计算结果中取最后的有效token的预测作为回复用户的首token进行自回归计算,当前hello world例子中选择world所生成的token作为首token,放入模型开头开始计算,此时的维度为[1,1,3072]。每一次onnx计算完取下一个token放置到开头接着计算,直到得到终止符号eos等。
总结
以上是大模型整体计算的完整流程。将这个树干打牢后想往哪进阶都可以。可以自行在huggingface中找不同的onnx来解析并分析其计算图来学习。
更多推荐

所有评论(0)