"""MLA:Multi-head Latent Attention(多头潜在注意力)。 这是 DeepSeek-V2 起最核心的省显存设计。 标准 MHA 的 KV cache 每 token 要存 2 * n_heads * head_dim 个数; MLA 先把每个 token 压成一个低维潜在向量 c_kv(kv_lora_rank 维), 用的时候再用 W_UK / W_UV 升维还原出 K 和 V。于是 cache 每 token 只需要存 kv_lora_rank + qk_rope_head_dim 个数,nano 档就是 64+16=80,而同规模 MHA 需要 2*8*32=512,省了 6.4 倍。 一个麻烦:RoPE 是位置相关的旋转,如果 K 是从 c_kv 现算出来的, 旋转矩阵没法和 W_UK 交换位置,"矩阵吸收"技巧就失效了。 DeepSeek 的解法是「解耦 RoPE」:额外切出 qk_rope_head_dim 维专门承载位置信息, 这部分所有 head 共享、直接缓存;剩下的 qk_nope 部分完全不带位置信息。 两条实现路径: naive —— 把 K/V 显式还原出来,走 F.scaled_dot_product_attention(能吃到 flash 内核), 训练时用它最快。 absorb —— 把 W_UK 吸收进 Q、W_UV 吸收进输出投影,全程在潜在空间里算注意力, 长上下文解码时中间激活小得多。 两条路径数学上完全等价,tests.py 里有对拍。 """ import math from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F from .layers import RMSNorm, apply_rope class MLA(nn.Module): def __init__(self, cfg): super().__init__() self.dim = cfg.dim self.n_heads = cfg.n_heads self.q_lora_rank = cfg.q_lora_rank self.kv_lora_rank = cfg.kv_lora_rank self.qk_nope_head_dim = cfg.qk_nope_head_dim self.qk_rope_head_dim = cfg.qk_rope_head_dim self.qk_head_dim = cfg.qk_head_dim self.v_head_dim = cfg.v_head_dim self.attn_impl = cfg.attn_impl self.softmax_scale = self.qk_head_dim ** -0.5 # ---- Query 投影:可选低秩分解(大模型上能省不少参数)---- if self.q_lora_rank == 0: self.wq = nn.Linear(self.dim, self.n_heads * self.qk_head_dim, bias=False) else: self.wq_a = nn.Linear(self.dim, self.q_lora_rank, bias=False) self.q_norm = RMSNorm(self.q_lora_rank, cfg.norm_eps) self.wq_b = nn.Linear(self.q_lora_rank, self.n_heads * self.qk_head_dim, bias=False) # ---- KV 联合压缩:一次投影同时产出潜在向量 c_kv 和共享的 k_pe ---- self.wkv_a = nn.Linear(self.dim, self.kv_lora_rank + self.qk_rope_head_dim, bias=False) self.kv_norm = RMSNorm(self.kv_lora_rank, cfg.norm_eps) self.wkv_b = nn.Linear(self.kv_lora_rank, self.n_heads * (self.qk_nope_head_dim + self.v_head_dim), bias=False) self.wo = nn.Linear(self.n_heads * self.v_head_dim, self.dim, bias=False) self.dropout_p = cfg.dropout # 推理缓存(训练时为 None) self.kv_cache: Optional[torch.Tensor] = None self.pe_cache: Optional[torch.Tensor] = None # ------------------------------------------------------------------ def setup_cache(self, max_batch: int, max_seq_len: int, device, dtype): self.kv_cache = torch.zeros(max_batch, max_seq_len, self.kv_lora_rank, device=device, dtype=dtype) self.pe_cache = torch.zeros(max_batch, max_seq_len, self.qk_rope_head_dim, device=device, dtype=dtype) def clear_cache(self): self.kv_cache = None self.pe_cache = None def cache_bytes(self, seq_len: int) -> int: if self.kv_cache is None: return 0 per_token = self.kv_lora_rank + self.qk_rope_head_dim return per_token * seq_len * self.kv_cache.element_size() # ------------------------------------------------------------------ def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, start_pos: int = 0, mask: Optional[torch.Tensor] = None) -> torch.Tensor: B, T, _ = x.shape end = start_pos + T # ---------------- Query ---------------- if self.q_lora_rank == 0: q = self.wq(x) else: q = self.wq_b(self.q_norm(self.wq_a(x))) q = q.view(B, T, self.n_heads, self.qk_head_dim).transpose(1, 2) # (B,H,T,qk) q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) q_pe = apply_rope(q_pe, cos, sin) # ---------------- KV 压缩 ---------------- kv = self.wkv_a(x) # (B,T,c+rope) c_kv, k_pe = kv.split([self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) c_kv = self.kv_norm(c_kv) # 缓存的就是它 k_pe = apply_rope(k_pe.unsqueeze(1), cos, sin).squeeze(1) # (B,T,rope),所有 head 共享 if self.kv_cache is not None: self.kv_cache[:B, start_pos:end] = c_kv.to(self.kv_cache.dtype) self.pe_cache[:B, start_pos:end] = k_pe.to(self.pe_cache.dtype) c_kv_all = self.kv_cache[:B, :end].to(x.dtype) k_pe_all = self.pe_cache[:B, :end].to(x.dtype) else: c_kv_all, k_pe_all = c_kv, k_pe if self.attn_impl == "naive": out = self._attn_naive(q_nope, q_pe, c_kv_all, k_pe_all, mask) else: out = self._attn_absorb(q_nope, q_pe, c_kv_all, k_pe_all, mask) out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.v_head_dim) return self.wo(out) # ------------------------------------------------------------------ def _attn_naive(self, q_nope, q_pe, c_kv_all, k_pe_all, mask): """显式还原 K/V,交给 SDPA(可命中 flash-attention 内核)。""" B, H, T, _ = q_nope.shape S = c_kv_all.shape[1] kv = self.wkv_b(c_kv_all).view(B, S, H, self.qk_nope_head_dim + self.v_head_dim) kv = kv.transpose(1, 2) # (B,H,S,nope+v) k_nope, v = kv.split([self.qk_nope_head_dim, self.v_head_dim], dim=-1) k = torch.cat([k_nope, k_pe_all.unsqueeze(1).expand(B, H, S, -1)], dim=-1) q = torch.cat([q_nope, q_pe], dim=-1) return F.scaled_dot_product_attention( q, k, v, attn_mask=mask, is_causal=(mask is None and T > 1), dropout_p=self.dropout_p if self.training else 0.0, scale=self.softmax_scale, ) def _attn_absorb(self, q_nope, q_pe, c_kv_all, k_pe_all, mask): """矩阵吸收:全程在 kv_lora_rank 维的潜在空间里做注意力。""" B, H, T, _ = q_nope.shape W = self.wkv_b.weight.view(H, self.qk_nope_head_dim + self.v_head_dim, self.kv_lora_rank) W_uk = W[:, : self.qk_nope_head_dim] # (H, nope, c) W_uv = W[:, self.qk_nope_head_dim:] # (H, v, c) # q_nope @ W_UK —— 把升维矩阵吸收进 Q,K 就不用还原了 q_absorb = torch.einsum("bhtd,hdc->bhtc", q_nope, W_uk.to(q_nope.dtype)) scores = (torch.einsum("bhtc,bsc->bhts", q_absorb, c_kv_all) + torch.einsum("bhtd,bsd->bhts", q_pe, k_pe_all)) * self.softmax_scale S = c_kv_all.shape[1] if mask is not None: scores = scores + _to_additive(mask, scores.dtype) elif T > 1: causal = torch.ones(T, S, dtype=torch.bool, device=scores.device).tril(S - T) scores = scores.masked_fill(~causal, torch.finfo(scores.dtype).min) attn = scores.softmax(dim=-1, dtype=torch.float32).type_as(scores) if self.training and self.dropout_p > 0: attn = F.dropout(attn, self.dropout_p) x_lat = torch.einsum("bhts,bsc->bhtc", attn, c_kv_all) # 仍在潜在空间 return torch.einsum("bhtc,hdc->bhtd", x_lat, W_uv.to(x_lat.dtype)) # 输出时才用 W_UV 还原 def _to_additive(mask: torch.Tensor, dtype) -> torch.Tensor: if mask.dtype == torch.bool: return torch.zeros_like(mask, dtype=dtype).masked_fill(~mask, torch.finfo(dtype).min) return mask.to(dtype)