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^C与v_t^C; - RoPE 部分采用解耦 key
k_t^R; - 推理时每层、每 token 主要缓存
c_t^KV和k_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/4 或 1/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 · 场景切入
假设你在做一个电商客服大模型,用户和机器人聊了20轮,每轮都要把历史所有Key和Value存下来。现在用户量一上来,显存就爆了。你遇到过这种问题吗?有什么办法能压缩这个缓存?
- 问法 2 · 层层追问
大模型推理时你一般怎么处理历史token的Key和Value?……那如果序列很长,比如几十万token,显存怎么扛得住?……有没有办法在不损失太多质量的前提下,把缓存降到一个向量级别?
- 问法 3 · 直球架构
解释一下MLA的原理,它和标准MHA在KV缓存上有什么区别?再结合KV Cache说一下MLA为什么能节省显存,以及它在大模型长上下文推理中的优化意义。