为什么长文本推理总在显存崩溃边缘徘徊?深挖 KV Cache 的权衡逻辑
RuntimeError: CUDA out of memory 的元凶往往不是模型本身,而是那个在后台疯狂吞噬空间的 KV Cache。要理解大模型推理的成本,必须先拆解 KV Cache 到底在做什么。在 Transformer 的解码阶段,模型每生成一个新 token,都需要与之前所有已生成的 token 进行注意力计算。如果没有缓存机制,每产生一个新词,模型都要把之前所有的 token 重新计算一遍 Key 和 Value 向量,这种冗余计算的复杂度是 $O(n^2)$。这意味着当你生成到第 1000 个 token 时,为了算出这一个词,你得重复计算前面 999 个词的注意力权重。
KV Cache 的核心逻辑就是“空间换时间”。它将每一步计算出的 Key 和 Value 向量直接暂存在显存中,下次推理时直接调用,将重复计算量强行砍掉。但这带来了一个极其残酷的副作用:计算压力被完整地转嫁给了显存。
我们可以算一笔账:对于一个 7B 规模的模型,其隐藏层维度和头数决定了每个 token 产生的 KV Cache 体积。随着输入长度的增加,缓存占用量呈线性增长。在实际部署中,如果你在 A100 显卡上运行,可能会发现模型刚加载时显存占用很低,但随着对话轮数增加,显存占用会像阶梯一样迅速攀升,直到触发类似 Tried to allocate 2.5GB (GPU 0), but only 1.2GB is free 的 OOM 报错。这时候你增加 Batch Size 简直是自杀,因为每个并发请求都会独立占用一份巨大的 KV Cache 空间。
为了在推理性价比上寻找平衡,目前的工程实践主要集中在三个突破口。
首先是量化(Quantization)。传统的 KV Cache 采用 FP16 或 BF16 存储,这意味着每个元素占用 2 字节。通过将缓存量化到 INT8 甚至 INT4,可以理论上将显存占用直接减半甚至降低 75%。这在处理 32K 甚至 128K 超长文本时,是保证模型不崩溃的唯一手段。
其次是内存管理机制的变革,最典型的就是 vLLM 引入的 PagedAttention。传统的缓存分配是连续的,这会导致严重的内存碎片化(Internal Fragmentation),很多显存被预留了但没被使用。PagedAttention 借鉴了操作系统虚拟内存的分页思想,将 KV Cache 存储在不连续的内存块中,按需动态分配。这不仅解决了碎片问题,更让吞吐量得到了量级上的提升,使得在相同硬件上能支撑更多的并发请求。
最后是位置编码的优化。在处理超长文本时,RoPE(旋转位置编码)的外推能力至关重要。如果缓存过大且位置编码在长文本下失效,即便显存没爆,生成的文本也会出现严重的逻辑崩塌或重复。
总结来看,AI 推理的底层逻辑其实是一场关于“算力”与“显存”的博弈。KV Cache 将计算复杂度从平方级降到了线性级,但它让显存成为了推理性能的绝对瓶颈。在追求长文本能力的今天,优化 KV Cache 的存储效率,比单纯堆算力要关键得多。
