
论文复现工坊 No.4从零复现 DeepSeek-V2/V3 多头潜在注意力 MLA在现代开源大语言模型的演进中DeepSeek 系列提出的Multi-Head Latent Attention (MLA多头潜在注意力)是近年来在注意力机制上的最重大突破之一。传统的 Multi-Head Attention (MHA) 在长文本推理时 KV Cache 显存开销极大而 Grouped-Query Attention (GQA) 虽然压缩了 KV Cache但牺牲了一定的模型容量与多头表达能力。MLA 的核心创新在于通过低秩联合压缩Low-Rank KV Compression将 KV 缓存压缩为一个极小的潜在向量Latent Vector并在推理时巧妙利用矩阵吸收Matrix Absorption在完全不损失 MHA 表达能力的前提下将推理期 KV Cache 压缩了 93% 以上。本文给出 MLA 的核心数学推导与 PyTorch 纯张量实现。1. MLA 的低秩压缩与解耦 RoPE 原理MLA 对 Query、Key 和 Value 均进行了低秩压缩与解耦设计KV 联合压缩将原本巨大的 Key 和 Value 向量通过一个下投影矩阵 $W_{DKV}$ 压缩为低维潜在向量 $c_t^{KV} \in \mathbb{R}^{d_c}$其中 $d_c \ll d_{\text{model}}$解耦 RoPEDecoupled Rotary Position Embedding因为旋转位置编码 RoPE 具有位置相关性无法直接被静态低秩投影矩阵吸收。MLA 将 Query 和 Key 拆分为两个子空间一个携带 RoPE 位置信息的低维向量 $q_{i,t}^{R}, k_t^{R}$另一个是不带位置编码、纯由低秩潜在向量上投影生成的语义向量 $q_{i,t}^{C}, k_{i,t}^{C}$矩阵吸收技巧Matrix Absorption在自回归推理生成时KV Cache 只需要存储压缩后的潜在向量 $c_t^{KV}$ 和解耦位置键 $k_t^R$上投影矩阵 $W_{UK}$ 可以直接被吸收到 Query 的投影矩阵中完全不需要在显存中展开完整的 Key 和 Value 张量。$$\text{KV Cache 每 Token 占用} (d_c d_R) \times \text{Bytes}$$相比于标准 MHA 的 $2 \times n_{\text{heads}} \times d_h \times \text{Bytes}$显存占用降低了一个数量级。2. MLA 核心模块的 PyTorch 完整实现import torch import torch.nn as nn import torch.nn.functional as F import math class MultiHeadLatentAttention(nn.Module): def __init__( self, dim: int, n_heads: int 32, q_lora_rank: int 512, kv_lora_rank: int 512, qk_nope_head_dim: int 128, qk_rope_head_dim: int 64, v_head_dim: int 128 ): super().__init__() self.dim dim self.n_heads n_heads self.q_lora_rank q_lora_rank self.kv_lora_rank kv_lora_rank self.qk_nope_head_dim qk_nope_head_dim self.qk_rope_head_dim qk_rope_head_dim self.v_head_dim v_head_dim # 1. Query 低秩下投影与上投影 self.wq_a nn.Linear(dim, q_lora_rank, biasFalse) self.q_norm nn.RMSNorm(q_lora_rank) self.wq_b nn.Linear(q_lora_rank, n_heads * (qk_nope_head_dim qk_rope_head_dim), biasFalse) # 2. Key-Value 联合低秩压缩 self.wkv_a nn.Linear(dim, kv_lora_rank qk_rope_head_dim, biasFalse) self.kv_norm nn.RMSNorm(kv_lora_rank) self.wkv_b nn.Linear(kv_lora_rank, n_heads * (qk_nope_head_dim v_head_dim), biasFalse) # 3. 输出投影 self.wo nn.Linear(n_heads * v_head_dim, dim, biasFalse) self.scale 1.0 / math.sqrt(qk_nope_head_dim qk_rope_head_dim) def forward(self, x: torch.Tensor, freqs_cis: torch.Tensor None) - torch.Tensor: bsz, seqlen, _ x.shape # 1. 处理 Query 潜在压缩与拆分 q_latent self.q_norm(self.wq_a(x)) q self.wq_b(q_latent).view(bsz, seqlen, self.n_heads, self.qk_nope_head_dim self.qk_rope_head_dim) q_nope, q_rope torch.split(q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim-1) # 2. 处理 Key-Value 联合压缩与拆分 kv_combined self.wkv_a(x) kv_latent, k_rope torch.split(kv_combined, [self.kv_lora_rank, self.qk_rope_head_dim], dim-1) kv_latent self.kv_norm(kv_latent) # 上投影恢复多头 Key (No-PE 部分) 与 Value kv self.wkv_b(kv_latent).view(bsz, seqlen, self.n_heads, self.qk_nope_head_dim self.v_head_dim) k_nope, v torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim-1) # 3. 维度转置以适配多头计算 q_nope q_nope.transpose(1, 2) # (bsz, heads, seqlen, nope_dim) k_nope k_nope.transpose(1, 2) v v.transpose(1, 2) # (bsz, heads, seqlen, v_dim) # 4. 对带位置编码的 RoPE 部分进行单头/多头广播计算 (此处省略 apply_rope 步骤) q_rope q_rope.transpose(1, 2) k_rope k_rope.unsqueeze(1) # (bsz, 1, seqlen, rope_dim) 广播到所有 heads # 5. 拼接 No-PE 与 RoPE 计算打分矩阵 q_full torch.cat([q_nope, q_rope], dim-1) k_full torch.cat([k_nope, k_rope.expand(-1, self.n_heads, -1, -1)], dim-1) scores torch.matmul(q_full, k_full.transpose(-2, -1)) * self.scale probs F.softmax(scores.float(), dim-1).type_as(q) output torch.matmul(probs, v).transpose(1, 2).contiguous().view(bsz, seqlen, -1) return self.wo(output)3. 显存开销与理论性能对比在相同的模型维度Hidden Dim 4096, 32 Heads下对比单 Token 的 KV Cache 显存占用注意力机制架构单 Token KV 存储维度8k 上下文单并发 KV 显存相对 MHA 显存压缩率标准 MHA$2 \times 32 \times 128 8,192$128 MB1.00x (基准)GQA-8 (8 组 KV)$2 \times 8 \times 128 2,048$32 MB4.00x (节省 75%)MLA (DeepSeek 架构)$512 (\text{Latent}) 64 (\text{RoPE}) 576$9 MB14.22x (节省 93.0%)MLA 在保留了 32 个 Query 头完整表达能力的同时将单 Token 的 KV Cache 存储压缩到了仅 576 维使得单张 80GB 显卡在长文本下的并发吞吐能力获得了翻天覆地的跃升。4. 复现实战避坑建议RMSNorm 的缩放位置在 $W_{DKV}$ 投影后必须加上 RMSNorm 进行特征尺度约束否则低秩瓶颈层极易出现数值梯度下溢推理阶段的计算图优化在线上部署时切记不要在内存中显式通过 $W_{UK}$ 还原完整的 Key 矩阵而是将其与 Query 的全连接层进行预先乘法融合$Q_{\text{absorbed}} Q \cdot W_{UK}^T$直接计算 $Q_{\text{absorbed}} \cdot (c_t^{KV})^T$才能真正享受到低内存带宽带来的速度飞跃。