nanoVLM 中的 KV 缓存实现
Hugging Face 已在 nanoVLM 中从头实现了 KV(键值)缓存,这是一个用于训练视觉语言模型的简洁 PyTorch 代码库。此优化通过在自回归推理过程中消除冗余计算,使生成速度提升了 38%。
自回归生成中的计算冗余
自回归语言模型一次生成一个 token。在没有缓存的标准 transformer 实现中,模型必须处理整个序列——包括所有先前生成的 token——才能预测下一个 token。
由于 transformer 在内部是并行的,每次生成新 token 都需要对所有层进行一次完整的前向传播。这导致相对于序列长度,内存和计算需求呈二次增长。具体而言,模型在每一步都会重新计算所有先前 token 的键 (K) 和值 (V) 张量,即使这些 token 及其对应的投影并未改变。
KV 缓存如何优化推理
KV 缓存通过在处理完初始提示后存储每层计算得到的键和值,来缓解这种低效。模型不再重新处理整个序列,而是遵循以下增量工作流:
- 缓存初始状态: 首次前向后,每层计算得到的 $K$ 和 $V$ 被缓存。
- 增量计算: 在后续生成步骤中,模型仅计算最新 token 的 $K$ 和 $V$。
- 缓存更新: 将新的 $K$ 和 $V$ 追加到已有缓存中。
- 注意力计算: 当前 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
- OriginalKV Cache from scratch in nanoVLM