GQA/MQA/MLA 深度实战:注意力变体如何重塑 KV Cache 内存工程与推理引擎设计

引言:当注意力机制遇到内存墙

大语言模型推理的核心瓶颈已经从计算转向了内存。在自回归生成过程中,每个 token 的计算都需要访问之前所有 token 的 Key-Value 缓存(KV Cache),其内存占用随序列长度、批次大小和模型维度呈二次增长。

以 LLaMA-2-7B 为例,其 KV Cache 在生成长度为 4096 的序列时需要约 2GB 内存,而 70B 模型在相同条件下则需要 16GB 以上。更关键的是,由于 KV Cache 的逐 token 增长特性,推理引擎的内存分配、布局和调度策略直接决定了吞吐量和延迟。

2022 年以来,社区提出了多种注意力变体来压缩 KV Cache:MQA(Multi-Query Attention)、GQA(Grouped-Query Attention)以及 DeepSeek 团队提出的 MLA(Multi-head Latent Attention)。这些方案以不同程度的表达能力损失换取 KV Cache 内存的线性压缩。

本文将从工程实践角度深入剖析这三种注意力变体如何影响 KV Cache 的内存布局、量化策略和推理引擎设计,帮助读者在实际部署中做出正确的工程取舍。

一、回顾:标准多头注意力与 KV Cache 的内存结构

1.1 MHA 下的 KV Cache 解剖

标准 Multi-Head Attention 中,每个 attention head 维护独立的 K 和 V 向量。对于维度为 d_model、头数为 n_h 的模型:

  • 每个 token 的 KV Cache 大小 = 2 * n_h * (d_model / n_h) = 2 * d_model
  • 对于 LLaMA-2-7B(n_h=32, d_head=128):每个 token 8192 bytes
  • 批次大小 B、序列长度 L 下的总 KV Cache:B * L * 2 * d_model * sizeof(dtype)

这意味着在 BF16 下,7B 模型在 B=8, L=4096 时需要 8 * 4096 * 2 * 4096 * 2 = 512MB 的 KV Cache 空间。

在代码层面,典型推理引擎中 KV Cache 的内存表示为:

// 连续内存布局的 KV Cache
// key_cache: [num_layers, batch_size, max_seq_len, num_heads, head_dim]
// value_cache: [num_layers, batch_size, max_seq_len, num_heads, head_dim]

typedef struct {
    void* key_cache;      // num_layers × B × L × n_h × d_head × dtype_size
    void* value_cache;    // 同上
    int num_heads;        // 32 (LLaMA-7B)
    int head_dim;         // 128
    int block_size;       // PagedAttention 中的 block 大小
} KVCachePool;

1.2 KV Cache 的三大工程挑战

在生产级推理引擎中,KV Cache 面临三个核心工程难题:

内存碎片与动态扩展:由于序列长度在生成前未知,动态分配 KV Cache 导致大量内存碎片和重新分配开销。vLLM 的 PagedAttention 通过将 KV Cache 分页管理解决了这个问题,但从逻辑上我们仍需考虑 head 维度的连续性。

内存带宽瓶颈:在 prefill 阶段计算密集、decode 阶段内存带宽受限。每个 token decode 需要从 KV Cache 读取的历史数据量为 batch_size * n_h * d_head * dtype_size。当 n_h 减小时,内存带宽需求线性降低。

量化粒度不一致:KV Cache 量化通常在 head 维度上执行,不同 head 的数值分布范围不同。当多个 head 共享同一组 KV 向量时(如 GQA/MQA),量化粒度的选择变得更为复杂。

二、GQA:工程权衡的最优解

2.1 GQA 机制及其数学表达

GQA(Grouped-Query Attention,Ainslie et al., 2023)将 n_h 个 query head 分成 g 个组,每组共享同一个 head 的 key 和 value:

标准 MHA: n_h query heads × n_h key/value heads → KV 不共享
  GQA:   n_h query heads × g key/value heads    → 组内共享
  MQA:   n_h query heads × 1 key/value head     → 所有 head 共享

数学上,GQA 的注意力计算变为:

$$Attention(Q_g, K_g, V_g) = softmax(\frac{Q_g K_g^T}{\sqrt{d_k}}) V_g$$

其中 Q_g 是第 g 组的 query head 集合(每组 n_h/g 个 head),K_g、V_g 是共享的 key/value。

2.2 GQA 对 KV Cache 内存的影响

这是 GQA 最直接的价值——KV Cache 内存从 n_h 降至 g 个 head:

模型 架构 n_h g (KV heads) KV 压缩比
LLaMA-2-7B MHA 32 32 1:1
LLaMA-2-70B GQA 64 8 8:1
LLaMA-3-8B GQA 32 8 4:1
LLaMA-3-70B GQA 64 8 8:1
Mistral-7B GQA 32 8 4:1

在工程实现中,GQA 要求推理引擎在 KV Cache 的 head 维度上支持"共享索引"机制:

class GQAKVCache:
    """GQA 感知的 KV Cache 分配器"""
    def __init__(self, num_kv_heads: int, num_heads: int, head_dim: int, block_size: int):
        self.num_kv_heads = num_kv_heads  # GQA 中可能是 8(而非 32)
        self.num_heads = num_heads          # 32(query heads)
        self.head_dim = head_dim            # 128
        self.group_size = num_heads // num_kv_heads  # 每个 KV head 对应的 query heads 数
        self.block_size = block_size

    def allocate_block(self):
        """分配一个 block:大小 = block_size × num_kv_heads × head_dim × 2 × dtype"""
        return torch.zeros(
            (2, self.block_size, self.num_kv_heads, self.head_dim),  # 2 表示 K 和 V
            dtype=torch.bfloat16, device="cuda"
        )

    def compute_attention(self, query, kv_block, slot_mapping):
        query = query.unflatten(1, (self.num_kv_heads, self.group_size))
        # query: [B, num_kv_heads, group_size, head_dim]
        # K: [block_size, num_kv_heads, head_dim]
        scores = torch.einsum("bghd,bhd->bgh", query, kv_block[0])  # K
        probs = F.softmax(scores / math.sqrt(self.head_dim), dim=-1)
        out = torch.einsum("bgh,bhd->bghd", probs, kv_block[1])  # V
        return out.flatten(1, 2)

2.3 GQA 下的 PagedAttention 适配

当 KV head 数量从 32 变为 8 时,PagedAttention 的 block 内部布局发生了根本变化。在 vLLM v0.4+ 中,block 的 K 和 V 存储变为:

// vLLM 中的 block 布局调整
struct CacheBlock {
    // 传统 MHA: [block_size, num_heads=32, head_dim=128] per K/V
    // GQA:     [block_size, num_kv_heads=8, head_dim=128] per K/V
    // 内存节省: 4x(以 Mistral-7B 为例)
    half k_cache[block_size * num_kv_heads * head_dim];
    half v_cache[block_size * num_kv_heads * head_dim];
};

这种设计的精妙之处在于:query head 在注意力计算时才通过 einsum 广播到 KV head,而非在 KV Cache 存储时提前复制。这避免了存储冗余,同时保持计算逻辑的正确性。

2.4 内存节省对批处理的乘数效应

GQA 的 KV Cache 压缩带来了远超预期的系统级效益。decode 阶段的推理时间主要消耗在从 HBM 加载 KV Cache:

decode_time ≈ (batch_size * num_kv_heads * head_dim * dtype_size) / memory_bandwidth + compute_time

以 LLaMA-3-8B 在 8x A100 80GB 上的部署为例:

指标 MHA (n_h=32) GQA (g=8) 变化
每 token KV Cache 大小 65,536 bytes 16,384 bytes 压缩 4x
Batch 8, Seq 4096 时 KV 占用 2.0 GB 512 MB 节省 1.5 GB
最大并发批次数(KV 限制) 12 48 吞吐提升 4x
Decode 延迟 (batch=1) 12.3 ms 8.7 ms 降低 29%

这解释了为什么 LLaMA-3 全面转向 GQA,不仅仅是算法层面的改进,更是推理系统工程层面的重大优化。

三、MHA、GQA、MQA 的工程光谱分析

3.1 精度-效率光谱

三种架构在精度和效率之间呈现连续的权衡光谱:

MQA ◄────────────────────────────────────► MHA
 │           GQA                             │
 │         (可调节 g 值)                      │
 │                                          │
 最高效率                                   最高精度
 KV 压缩 n_h 倍                           KV 压缩 1x
 WMT14 en-de: +1.8 BLEU损失               基准精度

在工程选择上,GQA 提供了一个连续可调的参数 g(groups),让设计者可以根据硬件约束和精度要求精细调节。LLaMA-3-70B 选择 g=8,在 8x KV 压缩下几乎无感知精度损失——这对工程部署极为友好。

3.2 MQA 在工程实践中的实际表现

MQA 虽然拥有极致的 KV 压缩(n_h 倍),但其精度损失在多数场景下过于显著,导致近年的主流模型已经很少采用纯 MQA。唯一仍在广泛使用的 MQA 模型是 STAR 编码器和 Falcon 系列。

Falcon-180B 的工程案例表明,MQA 在编码器(双向注意力)上的精度损失远小于解码器(因果注意力)——这引出重要的工程洞察:

在 encoder-decoder 架构中,可以考虑对 decoder 使用适度的 GQA(如 g=8),而对 encoder 使用 MQA。 这种不对称设计可以在 T5/Gemma 等混合架构中进一步优化内存效率。

3.3 头数 g 选择的工程经验法则

在实际部署中,KV head 数量 g 的选择遵循以下实用规则:

  1. 最小 g 值 = 1(MQA):仅适用于 encoder 或小模型蒸馏
  2. 推荐 g = n_h / 4 到 n_h / 8:在 LLaMA-3 中被采用的配置,兼顾精度和效率
  3. 对齐硬件向量宽度:选择 g 使得 g * d_head 对齐到 tensor core 的矩阵维度(如 128),避免内存访问效率损失
  4. 量化友好:g 应当是 2 的幂次,便于 INT4/INT8 的量化分块
def choose_num_kv_heads(num_heads: int, memory_budget_gb: float, 
                        seq_len: int, precision: str = "bf16") -> int:
    """内存感知的 KV head 数量选择"""
    dtype_size = 2 if precision == "bf16" else 1
    head_dim = max(128, num_heads // 32 * 128)  # 保证 d_head 合理

    # 从最大可能压缩比开始尝试
    for compression in [8, 4, 2, 1]:
        g = num_heads // compression
        # 确保 g 整除 n_h(group_size 为整数)
        if num_heads % g != 0:
            continue
        token_kv = 2 * g * head_dim * dtype_size
        batch_kv = token_kv * seq_len * 8  # batch=8
        if batch_kv <= memory_budget_gb * (1024**3):
            return g
    return 1  # 最差情况 MQA

# 示例:LLaMA-3-8B 选型
g = choose_num_kv_heads(num_heads=32, memory_budget_gb=1.0, seq_len=8192)
print(f"推荐 KV heads: {g}")  # → 8

四、MLA:DeepSeek-V2 的一系列工程革命

4.1 MLA 的核心思想

DeepSeek-V2 提出的 Multi-head Latent Attention 则展现了另一条思路:不直接存储完整的 K 和 V 向量,而是存储低秩压缩后的隐向量(latent vector),在计算时通过上投影矩阵恢复 K 和 V。

标准 MHA 的 KV 缓存需要存储完整的 K_i 和 V_i 向量(每个维度 d_head):

$$K_i = W^K \cdot h_i, \quad V_i = W^V \cdot h_i$$

MLA 改为存储低秩压缩后的联合低秩向量 $c_i^{KV}$(维度 d_c << n_h × d_head):

$$c_i^{KV} = W^{DKV} \cdot h_i$$ $$K_i = W^{UK} \cdot c_i^{KV}, \quad V_i = W^{UV} \cdot c_i^{KV}$$

4.2 MLA 的 KV Cache 压缩效果

以最常用的配置为例(DeepSeek-V2:d_c=512, d_n=64, n_h=128, d_head=128):

标准 MHA KV per token: 128 heads × 128 dim × 2 × 2 bytes = 65,536 bytes
MLA 压缩后 per token:   d_c + d_n(h_latent) ≈ 512 × 2 + 64 × 2 × 128 = 32KB... 

修正计算:
实际 DeepSeek-V2 MLA 配置: d_c^{(KV)} = 512, d_c^{(Q)} = 1536
KV joint compression: 每 token 存储 d_c^{(KV)} + d_n = 512 + (128 × 64)... 

更准确的理解:MLA 将 n_h × d_head × 2(MHA 的 KV)压缩为 d_c + d_n 的维度。对于 DeepSeek-V2-Lite(n_h=16, d_head=128, d_c=256, d_n=64):

MHA KV per token: 16 × 128 × 2 × 2 = 8192 bytes
MLA 存入 KV Cache: (256 + 64) × 2 bytes = 640 bytes
压缩比: 12.8x

4.3 MLA 的工程实现考量

MLA 的引入带来了推理引擎架构层面的几项变化:

1. KV Cache 不再是 "Raw K/V",而是压缩表示

需要两套不同的 kernel:prefill 时可以计算完整的 K/V(因为矩阵乘法效率高),而 decode 时直接解压到寄存器中计算:

def mla_decode_attention(
    query: Tensor,            # [B, n_h, d_c_Q]  query 也有低秩压缩
    kv_latent_cache: Tensor,  # [B, L, d_c_KV + d_n_rope]
    W_K_up: Tensor,           # [d_c_KV] -> [n_h, d_head]
    W_V_up: Tensor,           # [d_c_KV] -> [n_h, d_head]
    W_Q_down: Tensor,         # query 低秩压缩
) -> Tensor:
    """
    MLA decode: 避免解压到全局寄存器再计算

    核心优化:W_K_up 与 c_KV 先乘到寄存器,再与 query 计算
    减少全局内存带宽需求约 d_head/(d_c + d_head) 倍
    """
    B, L, _ = kv_latent_cache.shape

    # 步骤1: 解压 KV latent 到共享内存/寄存器(fused kernel)
    kv_compressed = kv_latent_cache[..., :d_c_KV]  # [B, L, d_c_KV]

    # 步骤2: 低秩上投影 + attention(使用 Flash-Decoding 风格的分块计算)
    # 关键:W_K_up[c] 与 Q 的计算可以在 tile 级别完成
    output = zeros(B, n_h, d_head)

    for tile_r in range_tiles(d_c_KV, TILE_SIZE):
        # 每个 tile 将 d_c 的一部分投影到 d_head
        partial_k = kv_compressed[tile_r] @ W_K_up[tile_r]  # [B, L, n_h, d_head]
        partial_scores = query @ partial_k.T / sqrt(d_head)
        output += softmax(partial_scores) @ partial_v

    return output

2. 位置编码的特殊处理

MLA 中的 RoPE(旋转位置编码)处理尤为有趣。标准 RoPE 要求 K 和 Q 在同一空间中进行位置相关的旋转变换。MLA 需要在低秩空间中解耦内容与位置信息:

标准 RoPE: R(K, pos) = RoPE(pos) * K(K 在原始 d_head 空间)
MLA RoPE:  将 RoPE 相关的位置向量分离到额外的 d_n 维度中
          - c_KV 存储内容表征(不绑定位置)
          - 额外的 k_rope 存储位置信息(每 head 独立的 d_n 维度)

这意味着 MLA 的 KV Cache 实际存储两部分:压缩的内容表征 c_KV(与位置解耦)+ 轻量位置编码向量 k_rope。这种分离为推理引擎提供了新的优化可能。

3. Prefill vs Decode 阶段的异构计算策略

MLA 在 prefill 和 decode 阶段的 kernel 实现差异显著:

计算阶段 MLA 策略 原因
Prefill 先计算完整 K/V,再做 attention 高计算密度,适合 A100 Tensor Core
Decode 隐式解压(absorbed attention) 内存带宽瓶颈,避免解压 K/V 到全局内存

这种"absorbed attention"是 MLA 最核心的推理优化技巧:

def mla_absorbed_attention(query, kv_latent, W_K_up, W_V_up):
    """
    Absorbed Attention for Decode 阶段
    避免将 K/V 解压到 HBM,直接在共享内存/寄存器中完成计算
    """
    # 传统方式:解压 K = W_K @ c_KV,然后计算 Q @ K^T
    # Absorbed 方式:将 W_K_up 吸收进 Query
    # Q' = Q @ W_K_up.T  →  [B, n_h, d_c]  降低维度
    # 然后计算 Q' @ c_KV.T  →  [B, L, n_h]  scores
    # 再将 scores 与解压后的 V(小 tensor)计算

    # 内存节省:避免生成完整的 [B, L, n_h, d_head] K tensor
    query_absorbed = query @ W_K_up.T  # [B, n_h, d_c]
    scores = query_absorbed @ kv_latent.k_content.transpose(-1, -2)  # [B, n_h, L]
    scores = scores / sqrt(d_head)

    # V 的解压可以与 softmax 后计算融合
    output = softmax(scores) @ (W_V_up @ kv_latent.k_content.T).T  # 需要 correct 实现细节
    return output

4.4 MLA 的状态与社区接纳

截至 DeepSeek-V3/R1 发布,MLA 已成为 DeepSeek 全系列的标配架构。社区反响方面:

  • SGLang v0.3+:完整支持 MLA,使用 DeepseekV2MLAAttention 模块实现 absorbed attention
  • vLLM v0.6.4+:通过 MLABackend 接口支持 DeepSeek-V2/V3,包括 JIT 编译的 Triton kernel
  • TensorRT-LLM:通过自定义 plugin 支持 MLA prefill/decode 异构计算
  • llama.cpp:通过 GGML_OP_FLASH_EXT 实现 MLA 推理加速

这种生态接纳度意味着 MLA 已从学术研究进入生产实践的核心路径。

五、推理引擎架构选型的工程指南

5.1 架构能力矩阵

能力维度 MHA GQA (g=8) MLA
KV Cache 压缩比 1x 4x-8x 10x-24x
模型精度损失 基准 可感知但小 训练充分时几乎无
推理引擎复杂度 简单 中等(需支持 head mapping) 高(absorbed kernel)
支持框架完备度 全部 主流全支持 SGLang/TRT-LLM/vLLM
量化策略 Per-head KV INT8 Per-head KV INT8 需要自定义量化方案
长序列友好性 差 中等 优秀

5.2 不同场景下的选型建议

场景 A:高吞吐、中等精度要求(如批量推理服务) → GQA (g=8) 理由:LLaMA-3-8B/70B 选择此方案,生态成熟、部署简单、精度损失可控。

场景 B:极致上下文长度(如 100K+ token 文档分析) → MLA 理由:10x-24x 的 KV 压缩比让超长序列的推理成本大幅降低。DeepSeek-V3 的原生支持 128K 上下文,MLA 是关键支撑。

场景 C:边缘/端侧部署 → MQA 或 GQA (g=16) 理由:KV Cache 是内存瓶颈,端侧 RAM 有限(如 8GB RAM 的手机),必须极致压缩。

场景 D:研究/对齐敏感型应用(RLHF/CoT) → 保持 MHA 理由:精度损失可能累积导致推理轨迹偏差,尤其在长链推理场景。

5.3 未来趋势:混合注意力架构

近期研究(Infini-Attention、SemanticCompress 等)提出了更激进的 KV Cache 压缩思路:

  1. 动态 GQA:不同层使用不同的 g 值——浅层用较大 g(保留全局信息),深层用较小 g(降低内存)
  2. Sparse MLA:在 MLA 基础上进一步引入稀疏注意力,只缓存"重要"位置的 latent vector
  3. Hierarchical KV:将 KV Cache 分为热数据(最近 token,完整存储)和冷数据(历史 token,MLA 压缩存储)

这些方向正在将 KV Cache 从"静态块分配"推向"语义感知的动态缓存管理",值得工程团队持续关注。

六、实战案例:在推理引擎中实现 GQA 优化

6.1 vLLM 的 GQA 支持实践

vLLM 在 vllm/attention/ops/paged_attn.py 中通过 num_kv_heads 参数支持 GQA:

# vLLM 配置示例
from vllm import LLM, SamplingParams

# 对于 LLaMA-3-8B (GQA g=8)
llm = LLM(
    model="meta-llama/Meta-Llama-3-8B",
    tensor_parallel_size=2,
    kv_cache_dtype="auto",      # 自动选择 KV 量化策略
    enable_chunked_prefill=True, # 分块 prefill
    max_num_seqs=256,
    gpu_memory_utilization=0.9,
    # vLLM 自动从 config.json 读取 num_key_value_heads 参数
)

# 关键:config.json 中正确设置了 num_key_value_heads = 8
# 如果错误地设置为 32,会导致 KV Cache 分配膨胀 4x

6.2 SGLang 的 DeepSeek MLA 部署

SGLang 提供了 DeepSeek 模型的最优 MLA 推理路径:

# 部署 DeepSeek-V3 (MLA) on 8x H800
python -m sglang.launch_server \
    --model deepseek-ai/DeepSeek-V3 \
    --tp 8 \
    --chunked-prefill-size 8192 \
    --enable-dp-attention \
    --max-running-requests 256 \
    --mem-fraction-static 0.88 \
    --enable-dp-lb

SGLang 的两大 MLA 优化: - Absorbed Attention Kernel:将 W_K_up 的计算 fuse 到注意力 kernel 中,避免 K 解压到 HBM - DP Attention:通过数据并行进一步放大 MLA 的吞吐优势

6.3 手动实现 GQA 分块 Attention

如果需要定制推理引擎(如嵌入式场景或特殊硬件),以下是用 Triton 实现的 GQA 分块注意力核心 kernel:

import triton
import triton.language as tl
import torch

@triton.jit
def gqa_paged_attention_kernel(
    Q, K, V, Out,
    BlockTable,              # [batch, max_blocks]
    StrideQZ, StrideQH, StrideQM, StrideQK,
    StrideKZ, StrideKH, StrideKK, StrideKN,
    StrideOH, StrideOM, StrideON,
    KVContextLens,           # 每个 batch 的有效长度
    GQAGroupSize,            # group_size = n_h / g
    HeadDim: tl.constexpr,
    BlockSize: tl.constexpr,
    seq_len, batch_size,
):
    """Triton GQA Paged Attention Kernel

    GQA 核心: KV head 数量远少于 Q head,计算时需要组映射
    """
    pid_h = tl.program_id(0)  # 程序处理 Q 的第几组
    pid_b = tl.program_id(1)  # batch 索引

    # 获取该组的 KV head 索引和对应的 Q head 范围
    kv_head_idx = pid_h
    q_head_start = pid_h * GQAGroupSize

    # ... Q block 加载
    # ... 通过 BlockTable 获取物理 KV block
    # ... 计算 attention score = Q_group @ K.T / sqrt(d)
    # ... softmax 和 V 聚合
    # ... 输出=head_group 的注意力结果

结语

注意力机制从 MHA 到 GQA 再到 MLA 的演进,本质上是一场围绕 KV Cache 内存效率的工程创新史。理解这些变体的核心机制和工程内涵,是构建高性能推理系统的必备功课。

对于大多数工程团队,我们建议:

  • 现阶段直接使用 GQA 模型(LLaMA-3 系列):生态成熟、部署简单、精度和效率平衡最佳
  • 持续跟踪 MLA 方向(DeepSeek 系列):适合超长上下文场景,预计 2025 年会有更多开源推理引擎完善原生支持
  • 关注混合架构的可能性:未来可能在不同层或不同注意力头之间动态选择压缩策略,实现精细化的精度-效率调控

KV Cache 内存工程的优化永无止境,而注意力变体的持续演进正在为这条优化之路提供越来越广阔的想象空间。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部