部署多模态大模型时的 KV Cache:先算对,再优化
部署多模态大模型时的 KV Cache:先算对,再优化
站内搜索
直接问 AI

部署多模态大模型时的 KV Cache:先算对,再优化

引言:多模态大语言模型(MLLM)的部署困境

在部署多模态大语言模型(MLLM)如 LLaVA、Qwen-VL 或 InternVL 时,开发者往往会遇到一个致命的内存瓶颈:KV Cache 的显存爆炸。尤其是当我们处理高分辨率(如 4K)图像或长视频时,视觉 token 的数量随像素面积增长,序列很快就超出上下文窗口和预填充的计算预算。这篇先把 KV Cache 的显存算对,再谈优化的先后顺序。

KV Cache 显存消耗的数学推导

在 Transformer 架构的自回归解码阶段,为了避免重复计算先前 token 的 Keys 和 Values,我们会将它们缓存在显存中,这就是 KV Cache。每个 token 的 KV Cache 大小可以精准计算如下:

公式:$VRAM_{KV} = 2 \times layers \times kv\_heads \times head\_dim \times seq\_len \times bytes$

参数解析:

  • 2:分别代表 Key 和 Value。
  • layers:模型的解码器层数(例如 Llama-3-8B 为 32 层)。
  • kv_headsKV 头数,不是注意力头数。见下方说明。
  • head_dim:每个头的维度。
  • seq_len:当前序列长度,即 token 数量。
  • bytes:数据精度(如 FP16 占 2 bytes,INT8 占 1 byte)。

必须用 KV 头数,不是注意力头数

这是这个公式最容易算错的地方。现代模型普遍使用分组查询注意力(GQA):多个查询头共享一组 Key/Value 头。Llama-3-8B 有 32 个注意力头,但只有 8 个 KV 头。用注意力头数代入,结果会高估 4 倍。

seq, layers, head_dim, byt = 8000, 32, 128, 2
kv = lambda h: 2 * layers * h * head_dim * seq * byt / 2**30

kv(32)   # 3.91 GiB  —— 按注意力头算,错
kv(8)    # 0.98 GiB  —— 按 KV 头算,对

配置文件里对应的字段通常是 num_key_value_heads;若模型没有这一项(早期 MHA 架构),才等于 num_attention_heads

随序列长度是线性,不是二次方

公式里 seq_len 是一次项,所以KV Cache 显存随序列长度线性增长:token 翻倍,缓存翻倍。呈二次方增长的是注意力分数矩阵 $QK^T$,那是计算量不是常驻显存——而且用了 FlashAttention 之后它根本不会被完整materialise 出来。

把这两件事分清楚很重要,因为它们的优化手段完全不同:线性的缓存靠量化和分页管理,二次方的计算靠算子融合和减少 token 数。搞混了就会用错药。

高分辨率图像到底产生多少 token:
假设使用 ViT-L 提取视觉特征,patch 大小为 14×14。一张 1024×1024 的图像被切分为 $(1024/14)^2 \approx 5329$ 个视觉 token。若不做任何切分直接处理 4K 图像(3840×2160):

(3840 ÷ 14) × (2160 ÷ 14) = 274 × 154 ≈ 42,196 tokens

注意 token 数与像素面积成正比(与边长成二次方)。42,000 个 token 已经超过大多数模型的上下文窗口,这才是高分辨率必须切块的真正原因——不是显存装不下缓存(按上面的公式约 5 GiB,A100 装得下),而是序列根本放不进上下文窗口,且预填充阶段的注意力计算是这个长度的平方

AnyRes/动态分辨率:缓解视觉 Token 爆炸

为了处理高分辨率图像而不引起显存 OOM,业界广泛采用了 AnyRes(动态分辨率)图像切块(Tiling)策略。该策略将高清大图动态切分为多个低分辨率的局部块(Local Patches),并保留一张全局缩略图(Global Context)。


graph TD
    A[原始高分辨率图像 4K] --> B{AnyRes 动态切分}
    B --> C[全局视角 Global View 
调整至 336x336] B --> D[局部切块 Local Patches
切分为 NxM 个 336x336 块] C --> E(ViT 视觉编码器) D --> E E --> F[合并拼接 Feature Concat] F --> G[Projector
MLP 降维/压缩] G --> H[LLM 语言模型]

这种策略能够在捕捉局部高清细节的同时,保留全局语义上下文,同时通过 projector 或 pooling 层大幅压缩视觉 token 数量。

实战代码:使用 llama.cpp 部署 LLaVA 4-bit 量化推理

在边缘设备或消费级显卡上,我们可以通过 C++ 后端 `llama.cpp` 和 4-bit 量化方案(如 GGUF)来部署 MLLM,大幅减少显存和内存带宽的压力。


#!/bin/bash
# 克隆并编译 llama.cpp,支持 CUDA 硬件加速
git clone https://github.com/ggerganov/llama.cpp.git
cd llama.cpp
make LLAMA_CUDA=1

# 下载量化后的多模态模型 (LLaVA-1.5-7b) 与视觉投影仪 (mmproj)
wget https://huggingface.co/mys/ggml_llava-v1.5-7b/resolve/main/ggml-model-q4_k.gguf
wget https://huggingface.co/mys/ggml_llava-v1.5-7b/resolve/main/mmproj-model-f16.gguf

# 运行 llava-cli 进行图文多模态推理
./llava-cli \
  -m ggml-model-q4_k.gguf \
  --mmproj mmproj-model-f16.gguf \
  --image /path/to/high_res_input.jpg \
  -p "Describe the image in detail, paying attention to the intricate textures." \
  -c 4096 \
  -ngl 35 # 将部分层卸载至 GPU 显存

真正的瓶颈在预填充,不在解码

纯文本 LLM 的优化经验会把人引到错误的方向。文本对话的提示词通常几百个 token,绝大部分时间花在逐 token 解码上,所以优化重点是解码阶段的显存带宽。多模态不是这样。

视觉 token 全部出现在预填充阶段——它们是提示词的一部分,在生成第一个字之前就要一次性算完。于是成本结构完全反过来了:

  • 预填充:一次性处理 n 个 token,注意力计算量正比于 $n^2$。n = 5000 时是 2500 万次分数计算(每头每层),n = 42000 时是 17.6 亿次——增长 70 倍
  • 解码:每生成一个 token,只需读一遍 KV Cache,计算量正比于 $n$,显存正比于 $n$。

实际表现是:问一张高分辨率图片,第一个字迟迟不出来,一旦开始输出就很流畅。用户感知到的”慢”几乎全部来自首字延迟,而不是输出速度。如果你在优化输出速度却发现体感没改善,多半是找错了地方。

这个结论直接决定了优化的优先级:

  1. 先减少 token 数。这是唯一同时改善两端的手段——预填充按平方受益,缓存按线性受益。AnyRes 切块、池化、Token Merging 都属于这一类。把 42000 压到 3000,预填充计算量降到 1/196。
  2. 再压缩缓存精度。KV Cache 量化到 int8 可以把线性那一项减半,但对预填充的平方项毫无帮助。所以它解决的是”能不能同时服务更多请求”,不是”首字快不快”。
  3. 最后做分页管理。分页解决的是显存碎片和多请求共享,同样不改变单请求的计算量。

顺序搞反是很常见的:先上量化和分页,发现首字延迟一点没变,然后怀疑是硬件不够。先确认自己受限于哪一项,再选工具。判断方法很简单——分别记录首 token 延迟和后续 token 的平均间隔,看哪个占了总时间的大头。

关于显存本身的估算——包括权重、KV Cache 与固定开销三项如何拆分、多卡时 KV 头数必须能被卡数整除等约束——可以参考把三条已知的显存坑变成一个能用的估算器,那篇给出了可以直接套用的公式和它的失效边界。

生产环境陷阱与工程排雷 (Production Pitfalls)

1. FlashAttention 在视觉序列中的限制

虽然 FlashAttention-2 对 LLM 来说是降低 VRAM 的标配,但在 MLLM 中,视觉 token 通常密集地堆叠在提示词序列的开头。如果跨帧或大图的视觉 token 超出了 FlashAttention 的块大小(block size),依然可能出现 OOM。排雷建议:对连续的视觉 token 引入 PagedAttention (如 vLLM 架构),对长上下文进行分页显存管理,避免由于显存碎片化导致的溢出。

2. 视觉 Token 上下文截断 (Context Length Management)

许多开发者在多轮图文对话中,会将历史的每一帧图像的 token 全部拼接进 prompt。这使得 KV Cache 随着对话轮数迅速达到模型最大 context window(如 4K 或 8K)。排雷建议:不要缓存多轮的历史原始视觉特征。应将之前轮次的视觉理解结果总结为文本缓存,或者在 Projector 层之后引入 Token Merging (ToMe) 技术,丢弃冗余的背景视觉 token。

发表回复

向下探索