多头注意⼒(Multi-Head Attention, MHA) 是 Transformer 架构的核⼼组件,在⼤语⾔模型中扮演 着⾄关重要的⻆⾊。MHA 通过并⾏计算多个注意⼒头,能够让模型同时关注序列中不同位置的不同表 ⽰⼦空间,从⽽捕获更丰富的语义信息和位置关系。在标准 MHA 中,每个注意⼒头都维护独⽴的 Key (K) 和 Value (V) 缓存(KV Cache),这在推理阶段会带来显著的显存开销。对于⼀个 H 头的注意⼒ 层,KV Cache 的显存占⽤为 O(B × L × H × d),其中 B 是批次⼤⼩,L 是序列⻓度,d 是每个头的维 度。随着模型规模增⼤和上下⽂⻓度增加,KV Cache 成为部署⼤模型的主要显存瓶颈。
多头潜在注意⼒(Multi-Head Latent Attention, MLA) 是 DeepSeek-V2/V3 等模型提出的创新机 制,通过低秩分解显著优化 KV Cache 的显存占⽤。MLA 将所有注意⼒头的 KV 表⽰投影到⼀个共享的 低秩潜在空间(Compressed Key-Value dimension)(维度 Ckv),⽽不是为每个头存储独⽴的⾼ 维 KV,将显存占⽤降低⾄ O(B × L × Ckv),其中 Ckv << H × d,实现数倍显存节省。通过头依赖的 权重矩阵在查询时动态恢复各头的独⽴表⽰,MLA 在保持模型表达能⼒的同时⼤幅降低了显存开销。
MLA 的核⼼创新在于解耦内容匹配与位置建模。查询向量被分为两个独⽴部分:⼀部分处理语义内容 匹配,需要先通过头依赖的权重矩阵投影到共享的低秩空间再与缓存交互;另⼀部分处理位置关系, 直接与位置缓存进⾏相关性计算。两路径的相关性分数需要合理合并并进⾏适当缩放,同时⾃回归解 码需要确保模型只能看到当前及之前的位置,未来位置必须被屏蔽,最终通过归⼀化操作得到注意⼒ 分布。注意⼒加权后的结果位于低秩潜在空间中,需要通过头依赖的权重矩阵将其映射回各头的值空 间,恢复多头的独⽴表⽰能⼒。
以下代码实现了 MLA 的单步解码函数及演⽰⽤法。请阅读以下代码,并根据描述完成空缺部分。
import torch
import math
def mla_decode_step(
q_nope: torch.Tensor, # (B,H,Q,Dn)
q_pe: torch.Tensor, # (B,H,Q,Dr)
kv_cache: torch.Tensor, # (B,Lk,Ckv)
pe_cache: torch.Tensor, # (B,Lk,Dr)
wkv_b: torch.Tensor, # (H, Dn+V, Ckv)
softmax_scale: float,
causal_mask: torch.Tensor # (1,1,1,Lk);可⻅=1,不可⻅=0
):
# === 维度准备与权重切分 ===
Dn = q_nope.size(-1)
Vd = wkv_b.size(1) - Dn
k_nope_w = wkv_b[:, :Dn, :] # (H, Dn, Ckv)
v_w = wkv_b[:, -Vd:, :] # (H, Vd, Ckv)
# 1) 将 q_nope 投影到 KV 低秩空间
q_nope_proj = ____[1]____ # -> (B,H,Q,Ckv)
# 2) ⽤ kv_cache 做相关得到 scores_nope
scores_nope = ____[2]____ # -> (B,H,Q,Lk)
# 3) 合并两部分分数并缩放
scores_pe = torch.einsum("bhqr,btr->bhqt", q_pe, pe_cache)
logits = ____[3]____
# 4) 应⽤因果掩码(不可⻅位置置为 -inf)
logits = ____[4]____
attn = torch.softmax(logits, dim=-1)
# 5) 聚合得到输出并映射回值空间
x = torch.einsum("bhqt,btc->bhqc", attn, kv_cache) # (B,H,Q,Ckv)
out = ____[5]____ # (B,H,Q,Vd)
return out
def demo_mla_step():
B, H, Dn, Dr, Vd, Lk, Ckv = 2, 8, 64, 32, 128, 10, 512
q_nope = torch.randn(B, H, 1, Dn)
q_pe = torch.randn(B, H, 1, Dr)
kv_cache = torch.randn(B, Lk, Ckv)
pe_cache = torch.randn(B, Lk, Dr)
wkv_b = torch.randn(H, Dn + Vd, Ckv)
softmax_scale = 1.0 / math.sqrt(Dn + Dr)
causal_mask = torch.ones(1, 1, 1, Lk) # 全可⻅⽰例
out = mla_decode_step(q_nope, q_pe, kv_cache, pe_cache, wkv_b,
softmax_scale, causal_mask)
print("MLA step out shape:", out.shape) # torch.Size([2, 8, 1, 128])
if __name__ == "__main__":
demo_mla_step()