我们都知道,现在各个大模型在整体的网络结构上基本上并没有太大差异,都是基于 Transformer 中的 Decoder 改进而来,所以基本也就是对其中的模块或使用顺序进行了替换或优化,因此为了后面能够更好的理解 LLaMA 2 模型的原理,我们依旧会先逐一介绍其中各个小的模块的原理及实现,然后再来整体看 LLaMA 2 模型。
因此,在今天这篇文章中将要和大家介绍的便是大模型当中都会使用到的 Key-Value Cache ,简称 KV Cache。
关键字:大模型、LLaMA、DeepSeek、KV缓存、多头注意力机制、KV Cache
1. 动机#
由于在生成式模型(如 Transformer 解码器)的推理过程中,token 是逐时刻解码生成得到,每个时刻新生成 token 时都需要将先前所有已生成的 tokens 重新输入到模型中,所以这就会带来巨大的计算冗余和重复,使得解码过程越来越慢。
例如在解码生成第 $t$ 时刻时,会将第 $0$ 到 $t-1$ 时刻的 token 都输入到模型中进行编码;而在生成第 $t+1$ 时刻时,又会将第 $0$ 到 $t$ 时刻的 token 都输入到模型中进行编码,其中第 $0$ 到 $t-1$ 时刻就会涉及到重复计算的问题。
如图1所示,模型输入 prompt 的长度为 3,在对 $t=3$ 时刻解码时会将整个 prompt 作为输入然后解码得到一个长度为 3 的张量,并取最后一个时刻作为 $t=3$ 时刻的生成结果;在对 $t=4$ 时刻解码时,会将 $t=3$ 时刻的输入与输出拼接起来作为输入生成 $t=4$ 时刻的结果;后续以此类推。
可以看出整个过程中会涉及到大量重复和冗余的计算。
注意这里的两个关键词“重复”冗余“”,重复代表着需要去重,而冗余则是要去掉。
有人可能会问,那训练阶段为什么不会呢?
这因为在训练阶段是一次性将整个样本输入到模型中(通过注意力掩码 attention mask 来使得编码当前时刻注意力只能关注到过去的信息而不能看到未来的信息),且所有时刻的输出都将用于损失计算,所以不存在重复或无用计算的问题。
因此,为了避免这种重复和冗余的计算的工作,Key-Value 缓存应运而生。简单来说,它在每次的自注意力计算过程中,都会将此时计算得到的 Key 和 Value 进行缓存,后续在生成下一个时刻的 token 时,模型只需对当前时刻输入的一个 token 计算隐藏状态,再把缓存与该隐藏状态拼接即可得到完整的隐藏状态 Key 和 Value,这样就能显著加快推理速度,而这在长序列或交互式应用中优势非常明显。
2. 不采用 KV 缓存推理过程#
在介绍完 KV 缓存出现的动机以后,我们再来看 KV 缓存的原理到底是什么样的。
不过为了让大家更好的理解整个过程,我们先来带大家通过图示的方式来回顾一下没有 KV 缓存时的完整计算过程,以便稍后将两者进行对比,以便更容易理解。
2.1 计算原理图解#
下面,我们依旧以图1中的情境为例进行说明。
首先,我们需要明白的是当模型训练完成以后,在推理过程中各个权重参数是固定不变的,因此对于同样的输入部分,其输出结果也是不变的,这一前提我们要知道。
假定现在有一个训练完成的模型将用于推理任务,为了便于介绍我们只考虑涉及到 KV 缓存的部分,即自注意力机制的计算过程,如图2所示。
在图2所示的示例中,原始 prompt 输入为一个长度为 3 的序列,输入到模型已经首先完成 3 个线性变换分别计算得到 $Q$ 、$K$ 和 $V$ ,进一步完成注意力权重的计算已经最终的输出,即最上方的结果。这里需要注意的是,因为这是在推理阶段,所以在对输入序列进行编码时需要进行 attention mask 操作,以保证当前时刻不能看到未来信息。在得到第 $t=3$ 时刻的整个输出以后,模型将会取结果的最后一个时刻作为 $t=3$ 时候的生成结果。
此时,我们可以得到第一个结论:在对第 $t$ 个时刻解码时,生成内容中前 $t-1$ 个时刻的结果是无用的。
进一步,开始依次进行后续时刻的解码生成,如图3所示。
在图4中,模型对第 $t=4$ 时刻解码时,会将整个 prompt 以及到当前时刻为止已经生成的内容拼接起来作为输入来生成当前时刻的结果。同理,此时的输入首先将完成 3 个线性变换分别计算得到 $Q$ 、$K$ 和 $V$ ,然后进行后续计算。此时我们可以发现,在第 $t=4$ 时刻时,其输入序列的前 3 个 token (图中灰色部分)在第 $t=3$ 时刻中已经分别进行过了一次同 $W^Q$ 、 $W^K$ 和 $W^V$ 的线性变换,也就是说此时 $Q$ 、$K$ 和 $V$ 这3个矩阵的前三行都是之前已经计算过的(图中红框中的内容),这里是在重复计算。
进一步,得到 $Q$ 、$K$ 和 $V$ 后将完成自注意力计算过程并得到最后的输出。这里需要注意的是,因为此时已经是逐时刻进行解码,模型可以看到当前时刻之前的所有信息,所以不再需要进行 attention mask 操作。同时,我们已经可以知道,对于 $t=4$ 的输出结果,依旧只会取最后一个时刻对应的结果作为 $t=4$ 时刻生成的内容,其余时刻的结果无用。
当第 $t=5$ 时,其生成过程完全与上述过程一致,大家可以自行默想。
此时此刻相信大家已经发现了其中的门道:
① 如果我们每次在进行解码生成时,都将当前时刻计算得到的 $K$ 和 $V$ 进行缓存,那么在下一时刻进行解码时就不需要再重复计算对应的结果,这就是解决前面提到的“重复计算”的问题;
② 因为只取每个时刻输出结果的最后一个时刻作为当前时刻的生成结果,所以在①的基础上,当前时刻的输入仅使用上一时刻的输出的最后一个时刻即可,不需要和先前的输入拼接,这就是解决前面提到的“冗余计算”的问题;
③ 基于对于②的理解,所以我们不需要对 $Q$ 进行缓存,因为后面根本用不到。
上面这两点就是 KV 缓存的核心思想。下面我们先来通过一个示例来模拟一下不采用 KV 缓存推理时的整个过程。
2.2 从零开始实现#
首先,我们来快速实现多头注意力机制的计算过程,代码如下所示:
1 class Attention(nn.Module):
2 def __init__(self, embed_dim=64, num_heads=8):
3 super().__init__()
4 self.num_heads = num_heads
5 self.embed_dim = embed_dim
6 self.head_dim = embed_dim // num_heads
7 assert self.head_dim * num_heads == self.embed_dim
8 self.wq = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
9 self.wk = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
10 self.wv = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
11 self.wo = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
12
13 def forward(self, x: torch.Tensor, mask=None):
14 bsz, seq_len, _ = x.shape
15 xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)
16 # [ bsz, seq_len, embed_dim] @ [embed_dim, embed_dim] = [ bsz, seq_len, embed_dim]
17 xq = xq.view(bsz, seq_len, self.num_heads, self.head_dim) # [bsz, seq_len, num_heads, head_dim]
18 keys = xk.view(bsz, seq_len, self.num_heads, self.head_dim) # [bsz, seq_len, num_heads, head_dim]
19 values = xv.view(bsz, seq_len, self.num_heads, self.head_dim) # [bsz, seq_len, num_heads, head_dim]
20 print(f"keys(values) shape [bsz, seq_len, num_heads, head_dim]: {values.shape}")
21 xq = xq.transpose(1, 2) # [bsz, num_heads, seq_len, head_dim]
22 keys = keys.transpose(1, 2) # [bsz, num_heads, seq_len, head_dim]
23 values = values.transpose(1, 2) # [bsz, num_heads, seq_len, head_dim]
24 scores = torch.matmul(xq, keys.transpose(2, 3)) / math.sqrt(self.head_dim)
25 if mask is not None: # [1, 1, seq_len, seq_len]
26 scores = scores + mask
27 scores = F.softmax(scores.float(), dim=-1).type_as(xq)
28 output = torch.matmul(scores, values) # [bsz, num_heads, seq_len, head_dim]
29 output = output.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
30 return self.wo(output)上述代码就是多头注意力的完整实现过程,相关变量的形状也进行了清晰地标注,相信大家对这部分内容已经非常熟悉,在这里就不再赘述。如果有不熟悉的朋友,可以先阅读「 第3讲 Multi-Head Attention 代码实现教程:基于 PyTorch 手写多头注意力与掩码」内容。
2.3 推理过程模拟#
进一步,我们来完模拟整个解码过程,示例代码如下:
1 if __name__ == '__main__':
2 start_pos = seq_len = 3
3 bsz = 3
4 embed_dim, num_heads = 4, 2
5 total_len = 10
6 input_embeddings = torch.randn([bsz, seq_len, embed_dim])
7 attn = Attention(embed_dim=embed_dim, num_heads=num_heads, max_batch_size=10)
8 for cur_pos in range(start_pos, total_len):
9 print(f" =========== decoding at pos: {cur_pos} ============")
10 print(f"input_embeddings shape [bsz, seq_len, embed_dim]: {input_embeddings.shape}")
11 _, seqlen, _ = input_embeddings.shape
12 mask = None
13 if cur_pos == start_pos:
14 mask = torch.full((1, 1, seqlen, seqlen), float("-inf"))
15 mask = torch.triu(mask, diagonal=1)
16 print(mask)
17 print(f"mask shape [1, 1, seq_len, seq_len]: {mask.shape}]")
18 output = attn(input_embeddings, mask) # [bsz, seq_len, embed_dim]
19 print(f"attention output shape [bsz, seq_len, embed_dim]: {output.shape}")
20 next_token_hidden = output[:, -1].unsqueeze(1) # [bsz, 1, embed_dim]
21 print(f"next_token_hidden shape [bsz, 1, embed_dim] : {next_token_hidden.shape}")
22 input_embeddings = torch.cat([input_embeddings, next_token_hidden], dim=1)在上述代码中,第2行定义初始序列的长度,也就是 prompt 的长度, 其中 start_pos 表示从 $t=3$ 时刻开始解码生成。第 3~5 行是定义模型的相关参数。第6行是随机生成一个输入序列经过 embedding 后的结果,将作为初始输入。第7行是实例化一个多头注意力模块,这里我们将其简单的看成是解码器。第8行是开始逐时刻进行解码。第11~17行是构建对 prompt 编码时所需要的 attention mask 输入。第18、20行是得到当前时刻的解码输出,然后取最后一个时刻作为结果,并将其扩维到 [bsz, 1, embed_dim]。第22行则是将当前时刻的输入和结果拼接起来,作为下一个时刻的输入,并再次解码生成。
理解上述代码的时候,建议对照上面的图示过程。以上完整示例代码可参见 Code/C06_SelfAttention/C01_attention_no_kv_cache.py 文件。
2.4 输出结果分析#
在上述代码执行结束以后,将会得到类似如下结果:
1 =========== decoding at pos: 3 ============
2 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 4])
3 tensor([[[[0., -inf, -inf],
4 [0., 0., -inf],
5 [0., 0., 0.]]]])
6 mask shape [1, 1, seq_len, seq_len]: torch.Size([1, 1, 3, 3])]
7 keys(values) shape [bsz, seq_len, num_heads, head_dim]: torch.Size([3, 3, 2, 2])
8 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 4])
9 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
10
11 =========== decoding at pos: 4 ============
12 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 4, 4])
13 keys(values) shape [bsz, seq_len, num_heads, head_dim]: torch.Size([3, 4, 2, 2])
14 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 4, 4])
15 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
16
17 =========== decoding at pos: 5 ============
18 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 5, 4])
19 keys(values) shape [bsz, seq_len, num_heads, head_dim]: torch.Size([3, 5, 2, 2])
20 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 5, 4])
21 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
22
23 =========== decoding at pos: 6 ============
24 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 6, 4])
25 keys(values) shape [bsz, seq_len, num_heads, head_dim]: torch.Size([3, 6, 2, 2])
26 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 6, 4])
27 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
28
29 =========== decoding at pos: 7 ============
30 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 7, 4])
31 keys(values) shape [bsz, seq_len, num_heads, head_dim]: torch.Size([3, 7, 2, 2])
32 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 7, 4])
33 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
34
35 =========== decoding at pos: 8 ============
36 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 8, 4])
37 keys(values) shape [bsz, seq_len, num_heads, head_dim]: torch.Size([3, 8, 2, 2])
38 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 8, 4])
39 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
40
41 =========== decoding at pos: 9 ============
42 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 9, 4])
43 keys(values) shape [bsz, seq_len, num_heads, head_dim]: torch.Size([3, 9, 2, 2])
44 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 9, 4])
45 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])对于上述结果,大家可以自行进行分析,关键一点就是观察各个变量的输出形状。
看到这里,我们想要给大家抛出一个疑问,在真正的推理场景中不同用户同时向模型发起请求,那么同一个时间段内模型接收到的多个样本长度肯定是不一样的,那么此时将其作为一个 batch 输入到模型中进行解码生成应该如何处理呢?
换句话说,在一个 batch 中,prompt 的初始长度不同,如何根据 prompt 同时生成多个样本的输出内容?
大家可以先想一想,在介绍 LLaMA 2 完整实现过程时再来交这部分作业。
下面,我们正式进入到 KV 缓存的讲解。
3. 采用 KV 缓存推理过程#
3.1 计算原理图解#
在经过图2和图3的介绍以后我们已经对于为什么要使用 KV 缓存,以及如何使用 KV 缓存有了清晰的了解,下面我们直接通过图示来阐述整个过程。
如图4所示是采用 KV 缓存时,解码器对原始输入 prompt 的处理情况。可以看出,整个过程同图2所示一致,唯一区别在于此时将会把当前计算得到的 $K$ 和 $V$ 进行缓存,即右上角的 K cache 和 V cache 部分。
在对 prompt 编码并得到输出后,我们同样只取输出的最后一个时刻作为此时刻的生成结果,接下里开始生成后续时刻的内容
敲黑板! 注意,此时变化就来了。
如图5(左)所示,在解码第 $t=4$ 时刻时,当前时刻的输入便只有上一时刻的生成结果,其与 $W^Q$ 、 $W^K$ 和 $W^V$ 作用后便会分别得到一个形状为 $1\times4$ 的向量。此时,对于 $Q$ 来说保持不变,对于 $K$ 和 $V$ 来说将分别与第 $t=3$ 时候缓存得到的 K cache 和 V cache 拼接得到新的 $K$ 和 $V$ ,即图中虚线方框所示。
进一步,在本次自注意力计算完成以后将会得到第 $t=4$ 时刻的解码输出,即左上角粉色输出;同时,将会得到更新后的 K cache 和 V cache,即右上角部分。
以此类推,在解码第 $t=5$ 时刻时的过程将如图5(右)所示,大家可以看图自行理解。
下面,我们再来看如何基于上面第2.2节的代码,来进行改造,使得其支持 KV 缓存。
3.2 从零开始实现#
在基于上面的代码框架下,我们只需要再新定义两个变量 cache_k 和 cache_v 来缓存即可,核心示例代码如下:
1 class Attention(nn.Module):
2 def __init__(self, embed_dim=64, num_heads=8, max_batch_size=8, max_seq_len=100):
3 super().__init__()
4 ......
5 self.max_batch_size = max_batch_size
6 self.max_seq_len = max_seq_len
7 self.cache_k = torch.zeros((self.max_batch_size, self.max_seq_len, self.num_heads, self.head_dim))
8 self.cache_v = torch.zeros((self.max_batch_size, self.max_seq_len, self.num_heads, self.head_dim))
9
10 def forward(self, x: torch.Tensor, start_pos: int, mask=None):
11 ......
12 xq = xq.view(bsz, seq_len, self.num_heads, self.head_dim) # [bsz, seq_len, num_heads, head_dim]
13 xk = xk.view(bsz, seq_len, self.num_heads, self.head_dim) # [bsz, seq_len, num_heads, head_dim]
14 xv = xv.view(bsz, seq_len, self.num_heads, self.head_dim) # [bsz, seq_len, num_heads, head_dim]
15 self.cache_k[:bsz, start_pos: start_pos + seq_len] = xk
16 self.cache_v[:bsz, start_pos: start_pos + seq_len] = xv
17 keys = self.cache_k[:bsz, : start_pos + seq_len]
18 values = self.cache_v[:bsz, : start_pos + seq_len]
19 xq = xq.transpose(1, 2) # [bsz, num_heads, seq_len, head_dim]
20 keys = keys.transpose(1, 2) # [bsz, num_heads, cache_len + seq_len, head_dim]
21 values = values.transpose(1, 2) # [bsz, num_heads, cache_len + seq_len, head_dim]
22 ......在上述代码中,第2行新增了 max_batch_size 和 max_seq_len 这两个参数,分别表示最大的样本数量和最长样本长度,因为后面需要先预定义空的全零张量来缓存 $K$ 和 $V$。当然,不同的实现方式所需要的变量不同,例如在 huggingFace 的 Transformers 框架中,K cache 和 V cache 是通过一个叫做 past_key_values 元组缓存的,大概语句如下:
1 if past_key_value is not None:
2 key_states = torch.cat([past_key_value[0], key_states], dim=2)
3 value_states = torch.cat([past_key_value[1], value_states], dim=2)
4 past_key_value = (key_states, value_states) if use_cache else None 希望大家以后看到的时候不会感到陌生。
进一步,上面第7~8行代码便是预先定义的两个全为0的张量,后续用于缓存 K cache 和 V cache 这两个结果。第10行中多了一个 start_pos,用于指定后续应该将当前的 $K$ 和 $V$ 应该缓存到 cache_k 和 cache_v 中哪个时刻,即第15~16行对应的代码,其中 seq_len 表示当前输入序列的长度。
此时我们可以知道,对于 seq_len 来说第1次解码时为 prompt 的长度,后续则固定为 1 ,因为是逐时刻解码。第17~18行则是取对应完整的 $K$ 和 $V$ 用于后续自注意力机制的计算过程。
以上完整示例代码可参见 Code/C06_SelfAttention/C02_attention_kv_cache.py 文件。
3.3 推理过程模拟#
如果我们以图4和图5中的计算过程为例,在上述代码中解码第 $t=3$ 时刻时,此时 bsz = 1, start_pos = 0, seq_len=3,所以 cache_k[:bsz, start_pos: start_pos + seq_len] = xk 执行结束后,cache_k[:1, 0: 0 + 3] 中存放的就是图4中右上角的缓存。
进一步,在解码第 $t=4$ 时刻时, bsz = 1, start_pos = 3, seq_len=1,所以 cache_k[:0, 3: 3 + 1] = xk 执行结束后,cache_k[:1, : 3 + 1] 中存放的就是图5(左)中右上角的缓存。在解码第 $t=5$ 时刻时, bsz = 1, start_pos = 4, seq_len=1,所以 cache_k[:0, 4: 4 + 1] = xk 执行结束后,cache_k[:1, : 4 + 1] 中存放的就是图5(右)中右上角的缓存。
当然,上述整个过程我们可以通过如下代码进行模拟:
1 if __name__ == '__main__':
2 start_pos = seq_len = 3
3 bsz = 3
4 embed_dim = 4
5 total_len = 10
6 input_embeddings = torch.randn([bsz, seq_len, embed_dim])
7 prev_pos = 0
8 attn = Attention(embed_dim=embed_dim, num_heads=2, max_batch_size=10)
9 for cur_pos in range(start_pos, total_len):
10 print(f" =========== decoding at pos: {cur_pos} ============")
11 print(f"input_embeddings shape [bsz, seq_len, embed_dim]: {input_embeddings.shape}")
12 attn_input = input_embeddings[:, prev_pos:cur_pos]
13 bsz, seqlen, _ = attn_input.shape
14 mask = None
15 if seqlen > 1: # 只有在编码原始输入的时候才会用到 mask attention,用于遮蔽未来时刻的信息
16 mask = torch.full((1, 1, seqlen, seqlen), float("-inf"))
17 mask = torch.triu(mask, diagonal=prev_pos + 1)
18 print(mask)
19 print(f"attention input shape [bsz, seq_len, embed_dim]: {attn_input.shape}")
20 output = attn(attn_input, prev_pos, mask) # [bsz, seq_len, embed_dim]
21 print(f"attention output shape [bsz, seq_len, embed_dim]: {output.shape}")
22 next_token_hidden = output[:, -1].unsqueeze(1) # [bsz, 1, embed_dim]
23 print(f"next_token_hidden shape [bsz, 1, embed_dim] : {next_token_hidden.shape}")
24 input_embeddings = torch.cat([input_embeddings, next_token_hidden], dim=1)
25 prev_pos = cur_pos在上述代码中,第2~8行是定义相关超参数及实例化一个编码器,其中 prev_pos 用于指定当前解码的开始时刻,后续将传入编码器中作为内部的 start_pos 变量。第12行是取当前时刻解码器的输入,第1次时则取整个 prompt 作为输入。第14~18行是构造 attention mask,也只有在对 prompt 编码时才会用到。第20行是计算得到解码输出。第22行是取最后一个时刻作为当前时刻的生成结果,当然在第2次以后只会生成一个时刻的结果。第24行是构建得到整个 prompt 拼接每个时刻生成后的结果。
例如,在第2次执行上面这个 for 时,第12行中 attn_input = input_embeddings[:, 3:4],即只取上一个时刻生成的结果作为当前时刻的输入,此时 seqlen=1。
3.4 输出结果分析#
在上述代码执行结束以后,将会得到类似如下结果:
1 =========== decoding at pos: 3 ============
2 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 4])
3 tensor([[[[0., -inf, -inf],
4 [0., 0., -inf],
5 [0., 0., 0.]]]])
6 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 4])
7 keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: torch.Size([3, 3, 2, 2])
8 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 4])
9 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
10
11 =========== decoding at pos: 4 ============
12 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 4, 4])
13 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
14 keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: torch.Size([3, 4, 2, 2])
15 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
16 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
17
18 =========== decoding at pos: 5 ============
19 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 5, 4])
20 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
21 keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: torch.Size([3, 5, 2, 2])
22 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
23 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
24
25 =========== decoding at pos: 6 ============
26 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 6, 4])
27 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
28 keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: torch.Size([3, 6, 2, 2])
29 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
30 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
31
32 =========== decoding at pos: 7 ============
33 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 7, 4])
34 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
35 keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: torch.Size([3, 7, 2, 2])
36 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
37 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
38
39 =========== decoding at pos: 8 ============
40 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 8, 4])
41 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
42 keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: torch.Size([3, 8, 2, 2])
43 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
44 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])
45
46 =========== decoding at pos: 9 ============
47 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 9, 4])
48 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
49 keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: torch.Size([3, 9, 2, 2])
50 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 4])
51 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 4])在上述代码中,我们打印了每个时刻解码完成之后各个变量的输出形状,我们可以明显的看到,除了在解码 prompt 时,其余时刻attention input 序列长度都是1,keys(values) shape 拼接后的形状也符合我们的预期。
4. 张量并行下的 KV 缓存#
下面,我们再来简单看一下张量并行的环境下(即使用 ColumnParallelLinear 和 RowParallelLinear 实现的多头注意力机制),对于 KV 缓存的实现和使用,而这也是 LLaMA 系列模型所使用到的技术。
4.1 计算原理图解#
根据之前在文章「一张图让你看懂 LLaMA 多头注意力在GPU上的并行计算过程!从零实现附源码!」中的介绍可知,在实现张量并行时多头注意力的并行机制是以每个头为整体,将所有的头分配到不同的 GPU 设备上进行计算。
例如多头的数量是 4,一共有两个 GPU 设备,那么每个 GPU 设备将会同时计算两个头的自注意力计算过程,以此来完成整个并行过程。因此,我们便可以知道,在张量并行下的 KV 缓存其实就是在每个设备上同时进行多个头的计算并缓存各自的 Keys 和 Values,相互独立互不干扰。
如图6所示,便是 4 个头在两个 GPU 设备上使用带 KV 缓存的张量并行计算过程。对于每个设备来说均各自完成2个头的计算过程,同时分别拥有各自的 KV 缓存。
4.2 从零开始实现#
在清楚张量并行下的 KV 缓存的原理以后,我们再来基于第3.2小节中的代码继续改进以使得其支持张量并行,核心代码如下所示:
1 class Attention(nn.Module):
2 def __init__(self, args: ModelArgs):
3 super().__init__()
4 world_size = fs_init.get_model_parallel_world_size()
5 self.n_local_heads = args.n_heads // world_size
6 self.head_dim = args.dim // args.n_heads
7 self.wq = ColumnParallelLinear(args.dim, args.n_heads * self.head_dim,
8 bias=False, gather_output=False, init_method=lambda x: x, )
9 self.wk = ColumnParallelLinear(args.dim, args.n_heads * self.head_dim,
10 bias=False, gather_output=False, init_method=lambda x: x, )
11 self.wv = ColumnParallelLinear(args.dim, args.n_heads * self.head_dim,
12 bias=False, gather_output=False, init_method=lambda x: x, )
13 self.wo = RowParallelLinear(args.n_heads * self.head_dim, args.dim,
14 bias=False, input_is_parallel=True, init_method=lambda x: x, )
15 self.cache_k = torch.zeros((args.max_batch_size, args.max_seq_len,
16 self.n_local_heads, self.head_dim)).to(args.device)
17 self.cache_v = torch.zeros((args.max_batch_size, args.max_seq_len,
18 self.n_local_heads, self.head_dim)).to(args.device)
19
20 def forward(self, x, start_pos: int, mask=None):
21 ......
22 xq = xq.view(bsz, seqlen, self.n_local_heads, self.head_dim)
23 xk = xk.view(bsz, seqlen, self.n_local_heads, self.head_dim)
24 xv = xv.view(bsz, seqlen, self.n_local_heads, self.head_dim)
25 self.cache_k[:bsz, start_pos: start_pos + seqlen] = xk
26 self.cache_v[:bsz, start_pos: start_pos + seqlen] = xv
27 keys = self.cache_k[:bsz, : start_pos + seqlen]
28 values = self.cache_v[:bsz, : start_pos + seqlen]
29 xq = xq.transpose(1, 2)
30 keys = keys.transpose(1, 2)
31 values = values.transpose(1, 2)
32 print(f" Rank: {rank} keys shape: {keys.shape}, {keys.device}")
33 print(f" Rank: {rank} self.cache_k shape: {self.cache_k.shape}, {self.cache_k.device}") 在上述代码中,第5行是计算每个设备上并行计算多头的个数,一定要能够整除。第7~14行则是对应初始化自注意力机制中的 4 个线性变换,具体原理可以参见文章「一张图让你看懂 LLaMA 多头注意力在GPU上的并行计算过程!从零实现附源码!」的介绍。第15~18行是分别为每个设备初始化一个用于缓存的 Keys 和 Values 的变量。例如 4 个头在两个设备上并行,那么每个设备上都会各有一个 cache_k 和 cache_v 变量。第20~33行逻辑则与之前相同这里就不再赘述。
总之,记住一个原则:在张量并行中,把每个设备上的多头计算过程看成是非并行情况下的多头计算过程即可,每个设备独立并行计算,互不干扰。
以上完整示例代码可参见 Code/C06_SelfAttention/C03_attention_kv_parallel.py 文件。
4.3 推理过程模拟#
在完成张量并行下的多头注意力计算 KV 缓存代码实现后,我们以4个头,两个 GPU 设备,模型维度为 8 时的情况来开始模拟整个推理过程。首先,我们写一个简单的推理函数,代码如下所示:
1 def inference(args):
2 start_pos = seq_len = 3
3 bsz = 3
4 embed_dim, num_heads = args.dim, args.n_heads
5 total_len = 6
6 input_embeddings = torch.randn([bsz, seq_len, embed_dim]).to(args.device)
7 prev_pos = 0
8 rank = fs_init.get_model_parallel_rank()
9 attn = Attention(args).to(args.device)
10
11 for cur_pos in range(start_pos, total_len):
12 print(f" Rank: {rank} ===== decoding at pos: {cur_pos} ====== device: {args.device}")
13 print(f" Rank: {rank} input_embeddings shape [bsz, seq_len, embed_dim]:"
14 f" {input_embeddings.shape}, {input_embeddings.device}")
15 attn_input = input_embeddings[:, prev_pos:cur_pos]
16 bsz, seqlen, _ = attn_input.shape
17 print(f" Rank: {rank} attention input shape [bsz, seq_len, embed_dim]:"
18 f" {attn_input.shape}, {attn_input.device} ")
19 output = attn(attn_input, prev_pos) # [bsz, seq_len, embed_dim]
20 print(f" Rank: {rank} attention output shape [bsz, seq_len, embed_dim]:"
21 f" {output.shape}, {output.device} " )
22 next_token_hidden = output[:, -1].unsqueeze(1) # [bsz, 1, embed_dim]
23 print(f" Rank: {rank} next_token_hidden shape [bsz, 1, embed_dim] : "
24 f"{next_token_hidden.shape}, {next_token_hidden.device} ")
25 input_embeddings = torch.cat([input_embeddings, next_token_hidden], dim=1)
26 prev_pos = cur_pos在上述代码中,整理逻辑还是同我们在第3.3小节中介绍的一致,只是新增了部分变量的输出信息,以便我们后续根据结果进行分析。
4.4 输出结果分析#
在上述代码执行结束以后,将会得到整个推理完成后的结果以及各个变量信息的打印输出。这里,我们首先来看对输入 prompt 的编码输出结果,如下所示:
1 Rank: 0 当前进程 PID: 15546, 父进程 PPID: 15540
2 Rank: 0 initializing model parallel with size 2
3 Rank: 1 当前进程 PID: 15547, 父进程 PPID: 15540
4 Rank: 1 initializing model parallel with size 2
5 > initializing model parallel with size
6 > initializing context parallel with size 2
7 > initializing ddp with size 1
8 > initializing pipeline with size 1
9 Rank: 0 =========== decoding at pos: 3 =========== device: cuda
10 Rank: 0 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 8]), cuda:0
11 tensor([[[[0., -inf, -inf],
12 [0., 0., -inf],
13 [0., 0., 0.]]]], device='cuda:0')
14 Rank: 0 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 8]), cuda:0
15 Rank: 0 keys shape: torch.Size([3, 2, 3, 2]), cuda:0
16 Rank: 0 self.cache_k shape: torch.Size([32, 16, 2, 2]), cuda:0
17 Rank: 0 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 8]), cuda:0
18 Rank: 0 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 8]), cuda:0
19
20 Rank: 1 =========== decoding at pos: 3 =========== device: cuda
21 Rank: 1 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 8]), cuda:1
22 tensor([[[[0., -inf, -inf],
23 [0., 0., -inf],
24 [0., 0., 0.]]]], device='cuda:1')
25 Rank: 1 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 8]), cuda:1
26 Rank: 1 keys shape: torch.Size([3, 2, 3, 2]), cuda:0
27 Rank: 1 self.cache_k shape: torch.Size([32, 16, 2, 2]), cuda:0
28 Rank: 1 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 3, 8]), cuda:1
29 Rank: 1 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 8]), cuda:1根据上面结果我们可以看出,两个设备 cuda:0 和 cuda:1 均有对应输出。输入 input_embeddings 的形状为 [3, 3, 8];keys 的形状为 [bsz, n_local_heads, seqlen, head_dim],即上面的 [3, 2, 3, 2]。attention output 的输出为 [3, 3, 8],但是我们只取最后一个时刻,所以是 [3, 1, 8] ,即 next_token_hidden。此时 self.cache_k[:,:3] 缓存便是当前时刻对应的 Key 状态。
当前时刻解码介绍后,便开始循环解码生成后续结果,如下所示:
31 Rank: 0 =========== decoding at pos: 4 =========== device: cuda
32 Rank: 0 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 4, 8]), cuda:0
33 Rank: 0 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:0
34 Rank: 0 keys shape: torch.Size([3, 2, 4, 2]), cuda:0
35 Rank: 0 self.cache_k shape: torch.Size([32, 16, 2, 2]), cuda:0
36 Rank: 0 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:0
37 Rank: 0 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 8]), cuda:0
38
39 Rank: 1 =========== decoding at pos: 4 =========== device: cuda
40 Rank: 1 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 4, 8]), cuda:1
41 Rank: 1 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:1
42 Rank: 1 keys shape: torch.Size([3, 2, 4, 2]), cuda:1
43 Rank: 1 self.cache_k shape: torch.Size([32, 16, 2, 2]), cuda:1
44 Rank: 1 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:1
45 Rank: 1 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 8]), cuda:1
46
47 Rank: 0 =========== decoding at pos: 5 =========== device: cuda
48 Rank: 0 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 5, 8]), cuda:0
49 Rank: 0 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:0
50 Rank: 0 keys shape: torch.Size([3, 2, 5, 2]), cuda:0
51 Rank: 0 self.cache_k shape: torch.Size([32, 16, 2, 2]), cuda:0
52 Rank: 0 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:0
53 Rank: 0 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 8]), cuda:0
54
55 Rank: 1 =========== decoding at pos: 5 =========== device: cuda
56 Rank: 1 input_embeddings shape [bsz, seq_len, embed_dim]: torch.Size([3, 5, 8]), cuda:1
57 Rank: 1 attention input shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:1
58 Rank: 1 keys shape: torch.Size([3, 2, 5, 2]), cuda:1
59 Rank: 1 self.cache_k shape: torch.Size([32, 16, 2, 2]), cuda:1
60 Rank: 1 attention output shape [bsz, seq_len, embed_dim]: torch.Size([3, 1, 8]), cuda:1
61 Rank: 1 next_token_hidden shape [bsz, 1, embed_dim] : torch.Size([3, 1, 8]), cuda:1对于上述逐时刻解码生成的结果来说,变化地方在于 attention input 和 attention output 的形状均固定为了 [bsz, 1, embed_dim],而 cache_k 也在逐时刻缓存当前时刻计算得到的 Key 状态。
最后,需要注意的是,由于是并行计算,终端打印出的各个变量的信息在不同 GPU 设备间是穿插(乱序) 的。上面的输出结果顺序是我们手动进行整理的。
引用#
[1] LLaMA-1: https://github.com/mlwithme/llama/tree/llama_v1