nowordsxiaomu's picture
Initial release: DeepSeek-Flash-Mini nano (15M MoE, MLA+MTP)
5e6d9f5 verified
Raw
History Blame Contribute Delete
8.23 kB
"""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)