推理服务 · 内存

KV 缓存

预填充一次,之后每步一个词元

模型
GPT-2 small
每词元
36 KiB 2 × 12 × 768 × 2 B
1,024 个词元时
36 MiB
图中的头
第 2 层 · 头 11 64 维,画出 8 维

全部步骤

  1. 一次一个词元

    GPT-2 一次写一个词元,每个新词元都要关注之前的每一个。那些词元的键和值一旦算出就不再改变,所以推理系统把它们存进 KV 缓存,而不是每一步都重算整个前缀。

    计算的 K、V 行数:n(n + 1) / 2 → n

  2. 预填充

    提示词一趟处理完:它的所有词元一起经过每一层,以矩阵乘法的形式计算,每一层都把它们的键和值写入缓存。这些是一个头的真实 GPT-2 数值(它 64 个数中的前 8 个)。

    对整个提示词计算 K, V = X · WK, X · WV

  3. 一个解码步

    每个解码步让一个词元经过模型。它算出自己的 q、k、v,把 k 和 v 追加到缓存,再对所有缓存的键做注意力:只有一行分数,而不是一个矩阵。真实的 GPT-2 数值,来自第 2 层、头 11;把鼠标停在格子上看看。

    q · Kᵀ → softmax → · V

  4. 缓存会有多大

    缓存为每个词元、每一层、每个键/值头保存两个向量(K 和 V)。它随正在服务的每条序列的每个词元增长,在长上下文时会超过模型权重本身;分组查询注意力、潜在注意力和分页都是在对付它。

    2 × 层数 × kv 头数 × dhead × 2 字节

  5. 计算受限与内存受限

    预填充每读一个字节要做很多次运算,所以它受算力限制。一个解码步为了产生一个词元要读每个权重和缓存,所以它受内存带宽限制;把许多序列放在一个批次里服务,才是用满 GPU 算力的办法。

    每字节 FLOPs ≈ 每读一个权重处理的词元数

代码

# prefill: the whole prompt in one pass; every layer keeps its K and V
k, v = ln_1(x) @ W_k + b_k, ln_1(x) @ W_v + b_v      # (prompt_len, 64) per head
cache[layer] = (k, v)
# decode: one token per step
q, k_new, v_new = (ln_1(x_t) @ W) .split(768)          # just the new token
K = torch.cat([cache_k, k_new]); V = torch.cat([cache_v, v_new])
w = F.softmax(q @ K.T / 8, dim=-1)                    # one row of scores
out = w @ V

延伸阅读