KV Cache
前置阅读 : 从0实现LLM
相关论文:
Efficient Memory Management for Large Language Model Serving with PagedAttention
为什么要缓存KV?
Transformer计算分为两个阶段:
- Prefill阶段:对输入的Prompt计算直到生成第一个Token,这段时间被称为TTFT。
- Decode阶段(自回归生成阶段):生成的首Token和一开始的Prompt合并,作为新的输入再次计算,循环往复,直到达到max token或者生成结束符号。
从上述过程,我们很容易可以得到:在decode阶段,会重复用到之前计算过的K、V。因此,通过缓存K、V,可以重复利用之前的计算结果,从而大大提升推理速度。
随之而来的问题是,如何高效管理KV Cache?
显存是一个很昂贵且稀有的空间,如何管理KV Cache才能充分利用显存空间?
连续内存分配
很直接的想法就是连续、静态分配,一个request想要缓存自己的KV内容,需要提前申请,说明自己需要多大的空间,然后我们将空间分给它。
1 | import torch |
这个想法很顺理成章,如果你想用公共空间,那就预约,给出你要用的大小和时间,我们就可以很好的进行安排。但是KV内容无法做到确定性,它是不确定大小(与生成长度相关)、不确定生命周期的存在。因此,这种死板的分配管理方式很容易造成大量的碎片,无论是内部还是外部。
如果学过操作系统,很容易会联想到操作系统对内存空间的管理发展史
Paged Attention
现有管理方式缺点:
- 存在内部和外部碎片,没有充分利用显存空间
- 无法共享,一些复杂的解码策略如:Beam search 可以共享kv cache,但是这种静态的连续的内存管理方式没有办法实现KV cache的共享。
- 上面两点导致推理速度受显存空间的约束,简单粗放的管理方式限制了batch decode并行的数
目。
Batch Decode
Batch Decode 的核心思想是“拼车”。系统同时收集 $N$ 个请求当前生成的 Token,组成一个 Batch(形状为 [N, 1]),然后一次性喂给模型。这样,底层计算就从矩阵-向量乘法(GEMV)变成了矩阵-矩阵乘法(GEMM)。模型权重只需要从显存搬运一次,就可以同时对这 $N$ 个 Token 进行计算。这极大地提高了“算术强度(Arithmetic Intensity,即每次显存读取对应的计算量)”,将算力利用率提升到了较高水平。
1 | import torch |
decode阶段的瓶颈在于访存速度,一个解决方式就是一次访存多次计算来均摊成本,但是在合并多个请求计算时,有以下问题:
- 每个请求需要的生成长度长短不一,如果一起进行计算,就要进行填充,会浪费GPU的计算资源。
- 不同的请求到达时间不同,最简单的解决方式就是等待,攒够一定数目的请求再进行计算,但是这样会让先来的等待时间很长。
其实有一个解决方式是降低粒度,即我们不再根据请求进行调度,而是根据一次迭代,或者说一个token的生成为最小单位进行调度。但是这种方式依旧存在很多问题:
- 批处理带来了巨大的KV Cache开销,如何高效管理显存
- decode具有不同的算法,对显存管理造成了困难
- 无法确定输入和输出的长度,无法进行准确、高效的调度
解决方式:向计算机操作系统中的虚拟内存学习,采用虚拟内存的思想,将显存空间分为大小相等的块,给每个请求分若干个块,这样就解决了外部碎片的问题,每个请求的显存分配情况通过块表进行记录、管理,通过更改块表可以轻松实现KV Cache块的共享。
思想


学过OS就很容易理解了,块表的作用就是把KV块的逻辑号映射到真实的存储地址,这样就可以实现逻辑上连续但是物理上不连续了。
