跳到正文

MLA 与 KV Cache 怎么用?

Multi-Head Latent Attention 原理及推理优化机制

原题:请解释MLA(Multi-Head Latent Attention)的基本原理及其在大模型架构中的应用;同时说明KV Cache(Key-Value Cache)的作用机制及其在推理过程中的优化意义。

模型架构 · 京东真题

回答与解析

先区分 KV Cache 与 MLA

自回归解码第 t 步只对当前 token 新算 Query、Key 和 Value,将新的 K/V 追加到 KV Cache,再用当前 Query 对历史与当前 K/V 做注意力;直接复用的是历史 token 的 Key/Value,避免每一步重复投影历史前缀。若每个注意力头维度为 d_h,MHA 有 n_h 个 KV 头,那么单层、单 token 的 K/V 缓存元素数约为 2 n_h d_h;长度为 L 时再乘 L。GQA 将 KV 头数降为 n_kv<n_h,MQA 令 n_kv=1。这里要明确“每 token”与“整段序列”,不能在同一列混用。

MLA 的核心机制

DeepSeek-V2 的 MLA 是训练时就写进模型结构的低秩联合压缩,不是推理后再压缩已有 KV:

  • 从隐藏状态得到压缩 latent:c_t^KV=W^DKV h_t
  • 通过上投影恢复用于内容注意力的 k_t^Cv_t^C
  • RoPE 部分采用解耦 key k_t^R
  • 推理时每层、每 token 主要缓存 c_t^KVk_t^R,缓存维度应写成 d_c+d_h^R,并乘层数、序列长度、batch 和字节数。

上投影可与 Query 侧计算做矩阵吸收以减少显式恢复,但具体 kernel 和精度实现要以代码为准。MLA、GQA、MQA 都在降低 KV Cache,不过前者压缩 KV 表示,后两者共享 KV 头。

工程判断

比较方案时应在同一模型质量、dtype、batch、上下文长度和硬件下报告 KV bytes/token/layer、最大并发、TTFT、TPOT 与吞吐。论文或某个模型上的压缩比不能写成所有 MLA 的固定 1/41/8。MLA 是端到端训练的架构选择,不能类比成事后有损压缩,也不能无证据断言会“压模糊关键轮次”。

口语版讲法(约4分钟)

  • 先解释 KV Cache 为什么存在
  • 用每 token 维度比较 MHA、GQA 和 MQA
  • 说明 MLA 的压缩 latent 与解耦 RoPE key
  • 区分架构级压缩和事后有损压缩
  • 给出上线时应测的显存与延迟指标

这道题我会先把两个概念拆开。KV Cache 是自回归解码的通用优化:生成第 t 个 token 时,历史 token 的 Key 和 Value 已经算过,就按层缓存起来,当前步只对新 token 计算 Query、Key 和 Value,把新的 K/V 追加进缓存,再用当前 Query 对历史与当前 K/V 做注意力。这样省掉的是重复投影和历史部分的重复计算,但缓存会随层数、batch 和上下文长度线性增长。

比较缓存量时必须统一口径。假设每个头维度是 d h,MHA 有 n h 个 KV 头,那么单层、单 token 的 K 和 V 一共大约是 2 乘 n h 乘 d h 个元素;整段长度为 L 才再乘 L。GQA 把 KV 头数降成 n kv,让多个 Query 头共享一组 KV;MQA 更进一步,只保留一个 KV 头。所以它们节省的是“头数”这一维。不能把表头写成每 token,公式里却又偷偷乘序列长度。

MLA 的思路不同。以 DeepSeek-V2 公开的结构为例,它先把隐藏状态投影成一个较小的联合 KV latent,也就是 c t^KV,再通过上投影得到内容相关的 Key 和 Value。为了兼容 RoPE,它还保留了解耦的位置 Key 成分。推理时真正要缓存的不是只有一个 latent,而是 c t^KV 加上解耦的 RoPE key。单 token 缓存维度应写成 d c 加 d h^R,最后再乘层数、长度、batch 和数据类型字节数。部分投影还能通过矩阵吸收放到 Query 侧计算,但具体收益取决于实现。

这里有两个容易说错的边界。第一,MLA 是训练时就确定的注意力架构,不是把一个训练好的 MHA 模型在部署时随手做有损压缩;因此也不能没有实验就说它会把关键轮次压模糊。第二,论文里某个配置的压缩比例不是通用的四分之一或八分之一,比例由头数、latent 维度、RoPE 维度和精度共同决定。

如果做工程选型,我会在同模型质量、同上下文、同 batch、同硬件下比较每层每 token 的缓存字节数、最大并发、TTFT、TPOT 和吞吐。GQA、MQA 和 MLA 都能降低 KV 压力,但一个靠共享 KV 头,一个靠联合低秩表示。把公式口径和实测条件讲清楚,比背一个压缩倍数更重要。

再做一层正确性验收:用同一批前缀比较改造前后的 logits、固定种子生成和长上下文任务,确认误差来自数值精度还是架构差异;同时抓取每层缓存张量的真实 shape 和 dtype,与理论字节数逐项对账。若只看到显存下降,却没有质量回归、最大并发和端到端延迟,就还不能证明这个 MLA 实现可上线。

若模型采用张量并行,还要确认 latent 与 RoPE key 的分片轴、all-gather 和通信量;缓存元素少不代表通信自动更少。这些条件必须写进 benchmark 报告。

关键一句:MLA 推理时除压缩 KV latent 外,为 RoPE 还需缓存解耦的位置 Key 成分。

核验来源

  1. DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model
  2. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints
  3. Fast Transformer Decoding: One Write-Head is All You Need

面试官还可能这样问

  1. 问法 1 · 场景切入

    假设你在做一个电商客服大模型,用户和机器人聊了20轮,每轮都要把历史所有Key和Value存下来。现在用户量一上来,显存就爆了。你遇到过这种问题吗?有什么办法能压缩这个缓存?

  2. 问法 2 · 层层追问

    大模型推理时你一般怎么处理历史token的Key和Value?……那如果序列很长,比如几十万token,显存怎么扛得住?……有没有办法在不损失太多质量的前提下,把缓存降到一个向量级别?

  3. 问法 3 · 直球架构

    解释一下MLA的原理,它和标准MHA在KV缓存上有什么区别?再结合KV Cache说一下MLA为什么能节省显存,以及它在大模型长上下文推理中的优化意义。

同模块相关题目