nanoVLM 中的 KV 缓存实现

Hugging Face 已在 nanoVLM 中从头实现了 KV(键值)缓存,这是一个用于训练视觉语言模型的简洁 PyTorch 代码库。此优化通过在自回归推理过程中消除冗余计算,使生成速度提升了 38%

自回归生成中的计算冗余

自回归语言模型一次生成一个 token。在没有缓存的标准 transformer 实现中,模型必须处理整个序列——包括所有先前生成的 token——才能预测下一个 token。

由于 transformer 在内部是并行的,每次生成新 token 都需要对所有层进行一次完整的前向传播。这导致相对于序列长度,内存和计算需求呈二次增长。具体而言,模型在每一步都会重新计算所有先前 token 的键 (K) 和值 (V) 张量,即使这些 token 及其对应的投影并未改变。

KV 缓存如何优化推理

KV 缓存通过在处理完初始提示后存储每层计算得到的键和值,来缓解这种低效。模型不再重新处理整个序列,而是遵循以下增量工作流:

  1. 缓存初始状态: 首次前向后,每层计算得到的 $K$ 和 $V$ 被缓存。
  2. 增量计算: 在后续生成步骤中,模型仅计算最新 token 的 $K$ 和 $V$。
  3. 缓存更新: 将新的 $K$ 和 $V$ 追加到已有缓存中。
  4. 注意力计算: 当前 token 的查询 ($Q$) 与缓存的 $K$ 和 $V$ 结合,用以生成输出。

在实际使用中,这个缓存以每层字典的形式维护,包含形状为 (batch_size, num_heads, seq_len_cached, head_dim) 的 “key” 与 “value” 张量。

nanoVLM 中的技术实现

nanoVLM 中的实现涉及对三个主要组件的修改,以实现从全序列重新计算到增量更新系统的转变。

1. 注意力块更新

LanguageModelGroupedAttention 类中,forward 函数被修改为接受 block_kv_cache 参数。如果缓存已存在(表明模型不在预填充阶段),模型会为当前 token 计算 $K_{new}$ 和 $V_{new}$ 并将其与缓存的张量拼接。如果缓存不存在,则对提示进行初始计算。

2. 层级缓存追踪

LanguageModel 类现在实现了层级缓存追踪。它使用 start_pos 参数,以确保旋转位置编码与当前生成索引正确对齐,使模型能够知道新生成 token 相对于序列的绝对位置。

3. 生成循环的分支

VisionLanguageModel 中的 generate() 方法被拆分为两个独立阶段:

  • 预填充阶段: 模型对完整输入提示进行编码,并为所有层构建初始 KV 缓存。
  • 解码阶段: 模型顺序生成 token,使用缓存的键和值,以避免重新处理提示和先前生成的 token。

架构变更概览

模块 原始行为 新行为
LanguageModelGroupedAttention.forward 在每一步重新计算 $Q$, $K$, $V$ 使用并更新 KV 缓存
LanguageModel.forward 没有先前状态的记忆 追踪每层 KV 缓存,处理 start_pos
VisionLanguageModel.generate 单阶段生成循环 拆分为 预填充解码 阶段

权衡与影响

KV 缓存将每个 token 的推理复杂度从二次降低到 $O(\text{seq len})$,从而实现更快的推理并能够在消费级硬件上运行大型模型。然而,这种效率伴随一定的权衡:它会增加用于存储缓存的内存占用,并提升代码复杂度。此外,它可能限制某些推理方案,例如束搜索,因为后者可能需要更复杂的缓存管理。

Sources