MoE 的核心思想是“条件计算”——每个 token 只激活少数最相关的专家,大幅减少计算量。
密集 FFN 的瓶颈:每层 0.8B 参数,64 层 = 51B,每个 token 必须完整计算全部参数。
MoE 的解法:把 0.8B 的单一路径复制 256 份变成 138B 总参数,但 Router 只选 8 个最相关的专家(3.1%)干活。
| 部件 | 参数量 | 占比(70B 计) |
|---|---|---|
| Token 嵌入 | 128k × 8192 ≈ 1.05B | ~1.5% |
| Block(每层) | ≈ 1.1B | ~1.6% / 层 |
| 输出层(共享) | 0(共享嵌入权重) | 0% |
| 部件 | 参数量 | 计算式 |
|---|---|---|
| 多头注意力(整体) | 268M(0.27B) | 4 × d_model² |
| ├─ Q 投影 | 67M | d_model² |
| ├─ K 投影 | 67M | d_model² |
| ├─ V 投影 | 67M | d_model² |
| └─ O 投影 | 67M | d_model² |
| FFN (SwiGLU) 整体 | 805M(0.8B) | 3 × d_model × d_ff |
| ├─ 上投影 (up) | 268M | d_model × d_ff = 8192×32768 |
| ├─ 门控投影 (gate) | 268M | d_model × d_ff = 8192×32768 |
| └─ 下投影 (down) | 268M | d_ff × d_model = 32768×8192 |
| 单层 Block 总计 | ≈ 1.1B | 0.27B + 0.8B |
| 指标 | 密集 FFN (图B) | MoE (图C, E=256, K=8) | 说明 |
|---|---|---|---|
| 每层总参数 | 0.8B | 138B | MoE 总参数大 170 倍 |
| 每层激活参数 | 0.8B | 4.3B | MoE 激活参数仅多 5.4 倍 |
| 激活/总参比例 | 100% | 3.1% | ⭐ 关键:绝大部分参数闲置 |
| 每 token 计算量 | 0.8B FLOPs | 4.3B FLOPs | MoE 计算量增加有限 |
| 参数量/计算量比 | 1:1 | 32:1 | ⭐ MoE 用参数换计算 |
| 部件 | 参数量 | 计算式 |
|---|---|---|
| Router | 8192 × 256 ≈ 2.1M | d_model × E |
| 单专家 FFN (SwiGLU) | 3 × 8192 × 32768 ≈ 0.8B | 3 × d_model × d_ff |
| 所有专家(E=256) | 256 × 0.8B ≈ 138B | E × 单专家 |
| 推理激活(Top-K=8) | 8 × 0.8B ≈ 4.3B / 层 | K × 单专家 |
| 模型 | Block 层数 (N) | 隐藏维度 | 注意力机制 | FFN / MoE |
|---|---|---|---|---|
| DeepSeek-V4 | ~64 | ~8192 | CSA+HCA+SWA | MoE (256专家, 激活8) |
| MiniMax-M3 | ~56 | ~6144 | MSA(稀疏注意力) | MoE (激活23B) |
| GLM-5.2 | ~60 | ~7168 | DSA + IndexShare | MoE (激活40B) |
| Qwen3.8 | 64 | 5120 | Gated DeltaNet + Gated Attn | MoE (95B激活) / Dense |
| Kimi K3 | ~96 | ~10240 | KDA + Gated MLA | MoE (896专家, 激活16) |
| 机制 | 图 B 中替换位置 | 核心差异 | 2026 代表模型 |
|---|---|---|---|
| MHA | 图 B 蓝色框(标准实现) | 每个 Q 头独立 K/V(H 组),缓存 L·H·d | 早期模型 |
| MQA | 图 B 蓝色框(K/V 改为 1 头) | 所有 Q 共享单 K/V(1 组),缓存 L·1·d | 已淘汰 |
| GQA | 图 B 蓝色框(K/V 改为 G 头) | 分组共享 K/V(G 组),缓存 L·G·d | Qwen3.8 / GLM-5.2 / MiniMax-M3 |
| KDA | 图 B 蓝色框(完全替换为线性) | 无 Q/K/V 投影,循环状态,O(1) 显存 | Kimi K3 |
| MLA | 图 B 蓝色框(K/V 前加压缩层) | 缓存 Latent,解压后计算,L·d_latent | DeepSeek-V4 |
| 模型 | 发布时间 | 总参数 / 激活 | 架构 | 注意力机制 | 上下文长度 |
|---|---|---|---|---|---|
| DeepSeek-V4 | 2026.04 | 1.6T / 49B (MoE) | MoE | CSA+HCA+SWA | 1M |
| MiniMax-M3 | 2026.06 | 428B / 23B (MoE) | MoE + MSA | MSA | 1M |
| GLM-5.2 | 2026.06 | 744B / 40B (MoE) | MoE | DSA + IndexShare | 1M |
| Qwen3.8 | 2026.08 | 2.4T / 95B (MoE) | MoE | Gated DeltaNet + Gated Attn | 1M / 256K |
| 参数 | 取值 | 含义 |
|---|---|---|
| L | 4096 | 上下文长度(一次推理的 token 总数) |
| H | 32 | Q 头数 |
| d | 128 | 每头维度 |
| d_model | 4096 | = H × d,隐藏维度 |
| G | 4 | GQA 组数(每组 8 个 Q 头) |
| d_latent | 256 | MLA 潜在维度 |
| 精度 | FP16 | 缓存按 2 字节 / 元素计 |
| 机制 | KV 缓存大小 | 推理速度 | 模型质量 | 2026 代表模型 | 核心思想 |
|---|---|---|---|---|---|
| MHA | L·H·d | 🐢 较慢 | ⭐⭐⭐⭐⭐ | (早期模型) | 每 Q 独立 K/V |
| MQA | L·1·d | 🚀 极快 | ⭐⭐⭐ | (已淘汰) | 所有 Q 共享单 K/V |
| GQA | L·G·d | ⚡ 较快 | ⭐⭐⭐⭐ | Qwen3.8 GLM-5.2 MiniMax-M3 | 分组共享(G组) |
| KDA | 固定状态 (非KV缓存) | ⚡ O(L) | ⭐⭐⭐⭐ | Kimi K3 | 线性注意力+循环状态 |
| MLA | L·d_latent (d_latent << H·d) | 💾 省显存 | ⭐⭐⭐⭐ | DeepSeek-V4 | 低秩压缩 K/V |
class MHA(nn.Module): # 对应 图MHA-1 的蓝色区域 def __init__(self, d_model=4096, n_heads=32): self.n_heads = n_heads self.head_dim = d_model // n_heads # 128 self.W_q = nn.Linear(d_model, d_model) # 图MHA-2 的 Q 投影 self.W_k = nn.Linear(d_model, d_model) # 图MHA-2 的 K 投影(同样 32 头,浪费点 1) self.W_v = nn.Linear(d_model, d_model) # 图MHA-2 的 V 投影(同样 32 头,浪费点 2) self.W_o = nn.Linear(d_model, d_model) # 图MHA-1 的 O 投影 def forward(self, x, past_kv=None): B, T, _ = x.shape # 投影:每个 token 生成自己的一套 H=32 组 Q/K/V Q = self.W_q(x).view(B,T,self.n_heads,self.head_dim).transpose(1,2) # (B,32,T,128) K = self.W_k(x).view(B,T,self.n_heads,self.head_dim).transpose(1,2) # (B,32,T,128) V = self.W_v(x).view(B,T,self.n_heads,self.head_dim).transpose(1,2) # (B,32,T,128) # 图MHA-3 推理:把新 token 的 K/V 追加进缓存(缓存长度 = L,随 L 线性增长) if past_kv: K = torch.cat([past_kv[0], K], dim=2) # (B,32,L,128) V = torch.cat([past_kv[1], V], dim=2) # 图MHA-2 注意力:Q·K^T/√d → softmax → ·V(Q 扫过全部 L 个历史位置) scores = Q @ K.transpose(-2,-1) / math.sqrt(self.head_dim) # (B,32,T,L) attn = F.softmax(scores, dim=-1) out = attn @ V # (B,32,T,128) return self.W_o(out.transpose(1,2).reshape(B,T,-1)), (K, V) # O 投影 + 回传缓存
class MQA(nn.Module): # 对应 图MQA-1:K/V 改为 1 个头 def __init__(self, d_model=4096, n_heads=32): self.n_heads = n_heads self.head_dim = d_model // n_heads # 128 self.W_q = nn.Linear(d_model, d_model) # 图MQA-2 的 Q 投影(仍 32 头) self.W_k = nn.Linear(d_model, self.head_dim) # ⚡ 图MQA-2 的 K 投影改为只 1 个头 self.W_v = nn.Linear(d_model, self.head_dim) # ⚡ 图MQA-2 的 V 投影改为只 1 个头 self.W_o = nn.Linear(d_model, d_model) # O 投影 def forward(self, x, past_kv=None): B, T, _ = x.shape Q = self.W_q(x).view(B,T,self.n_heads,self.head_dim).transpose(1,2) # (B,32,T,128) K = self.W_k(x).unsqueeze(1) # (B,1,T,128) 只 1 个头 V = self.W_v(x).unsqueeze(1) if past_kv: K = torch.cat([past_kv[0], K], dim=2) # (B,1,L,128) 缓存极瘦 V = torch.cat([past_kv[1], V], dim=2) # 图MQA-2 广播:把 1 个头复制成 32 个头(等价 repeat_interleave) K = K.repeat(1, self.n_heads, 1, 1) # (B,32,L,128) V = V.repeat(1, self.n_heads, 1, 1) scores = Q @ K.transpose(-2,-1) / math.sqrt(self.head_dim) attn = F.softmax(scores, dim=-1) out = attn @ V return self.W_o(out.transpose(1,2).reshape(B,T,-1)), (K[:, :1], V[:, :1]) # 缓存只存 1 份
scores = Q·K^T 要求 Q 的头数 = K 的头数才能逐头点积——Q 有 32 个头,所以需要 32 份 K。于是把 4 份 K/V 各复制 H/G = 32/4 = 8 份,凑成 32 份,编号恰好和 32 个 Q 头一一对应(头 0-7 ↔ 组0,头 8-15 ↔ 组1,…)。repeat_interleave 就是“沿头维复制”这个操作;它是张量视图/广播,不会真的把缓存复制一份占额外显存——缓存里始终只有 G=4 组。
class GQA(nn.Module): # 对应 图GQA-1:K/V 改为 G 个头 def __init__(self, d_model=4096, n_heads=32, n_groups=4): self.n_heads, self.n_groups = n_heads, n_groups self.head_dim = d_model // n_heads # 128 self.W_q = nn.Linear(d_model, d_model) # 图GQA-2:Q 仍出 32 个头 self.W_k = nn.Linear(d_model, n_groups * self.head_dim) # ⚡ 图GQA-2:K 只出 G=4 个头 self.W_v = nn.Linear(d_model, n_groups * self.head_dim) # ⚡ 图GQA-2:V 只出 G=4 个头 self.W_o = nn.Linear(d_model, d_model) # O 投影 def forward(self, x, past_kv=None): B, T, _ = x.shape Q = self.W_q(x).view(B,T,self.n_heads,self.head_dim).transpose(1,2) # (B,32,T,128) K = self.W_k(x).view(B,T,self.n_groups,self.head_dim).transpose(1,2) # (B,4,T,128) V = self.W_v(x).view(B,T,self.n_groups,self.head_dim).transpose(1,2) # (B,4,T,128) # 图GQA-1 缓存:只追加 G=4 组,永远比 MHA 小 8 倍 if past_kv: K = torch.cat([past_kv[0], K], dim=2) # (B,4,L,128) V = torch.cat([past_kv[1], V], dim=2) # 图GQA-3 repeat_interleave:沿头维把 4 组复制成 32 组,与 Q 头对齐 K = K.repeat_interleave(self.n_heads // self.n_groups, dim=1) # (B,32,L,128) V = V.repeat_interleave(self.n_heads // self.n_groups, dim=1) # 视图操作,缓存仍是 4 组 # 之后的注意力与 MHA 完全一致 scores = Q @ K.transpose(-2,-1) / math.sqrt(self.head_dim) attn = F.softmax(scores, dim=-1) out = attn @ V # 回传缓存时只保留 G 组(复制前的原始 4 组) return self.W_o(out.transpose(1,2).reshape(B,T,-1)), (K[:, :self.n_groups], V[:, :self.n_groups])
W_kv 把它们压成低维 latent(潜在向量),缓存只存 latent,要用时再解压回来。因为 d_latent=256 << H·d=4096,缓存进一步缩小。
out = W_o · (Σ a_t · W_v_up · c_V_t) = (W_o · W_v_up) · Σ a_t·c_V_t。于是推理时完全不需要展开 K/V,全程只在 256 维 latent 上做注意力——省掉了每步「展开 32×128 的 K/V」的计算与显存搬运。
class MLA(nn.Module): # 对应 图MLA-1:K/V 改为 latent 压缩 def __init__(self, d_model=4096, n_heads=32, head_dim=128, d_latent=256): self.n_heads, self.head_dim, self.d_latent = n_heads, head_dim, d_latent self.W_q = nn.Linear(d_model, n_heads * head_dim) # 图MLA:Q 投影(同 MHA) self.W_kv = nn.Linear(d_model, 2 * d_latent) # ⚡ 图MLA-2 压缩层:x → (c_K, c_V) # 解压权重(训练用;推理被吸收进 Q/O,见 forward_fused) self.W_k_up = nn.Linear(d_latent, n_heads * head_dim) # 图MLA-2:c_K → 32×128 的 K self.W_v_up = nn.Linear(d_latent, n_heads * head_dim) # 图MLA-2:c_V → 32×128 的 V self.W_o = nn.Linear(n_heads * head_dim, d_model) # O 投影 def forward_naive(self, x, past_latent=None): # 训练/理解版:解压再算 B, T, _ = x.shape Q = self.W_q(x).view(B,T,self.n_heads,self.head_dim).transpose(1,2) # (B,32,T,128) c = self.W_kv(x) # (B,T,512) c_K, c_V = c[..., :self.d_latent], c[..., self.d_latent:] # 各 (B,T,256) if past_latent: # 图MLA-1 缓存:只存 latent c_K = torch.cat([past_latent[0], c_K], dim=1) # (B,L,256) c_V = torch.cat([past_latent[1], c_V], dim=1) K = self.W_k_up(c_K).view(B,-1,self.n_heads,self.head_dim).transpose(1,2) # (B,32,L,128) V = self.W_v_up(c_V).view(B,-1,self.n_heads,self.head_dim).transpose(1,2) scores = Q @ K.transpose(-2,-1) / math.sqrt(self.head_dim) # 与 MHA 相同 attn = F.softmax(scores, dim=-1) out = (attn @ V).transpose(1,2).reshape(B,T,-1) return self.W_o(out), (c_K, c_V) def forward_fused(self, x, past_latent=None): # 推理版:解压被吸收(图MLA-3) B, T, _ = x.shape Q = self.W_q(x).view(B,T,self.n_heads,self.head_dim) # (B,T,32,128) c = self.W_kv(x) c_K, c_V = c[..., :self.d_latent], c[..., self.d_latent:] if past_latent: c_K = torch.cat([past_latent[0], c_K], dim=1) # (B,L,256) c_V = torch.cat([past_latent[1], c_V], dim=1) # 吸收:q' = W_k_up^T · q(按头分块,W_k_up 每头占 (128,256)) Wk_up_T = self.W_k_up.weight.view(self.n_heads, self.head_dim, self.d_latent) q_abs = torch.einsum('bthd,hdL->bthL', Q, Wk_up_T) # (B,T,32,256) scores = torch.einsum('bthL,blL->bthl', q_abs, c_K) / math.sqrt(self.head_dim) attn = torch.softmax(scores, dim=-1) o_lat = torch.einsum('bthl,blL->bthL', attn, c_V) # 仍在 latent 空间 (B,T,32,256) # V 的解压也吸收进 O:out = W_o · (W_v_up^T · o_lat) Wv_up_T = self.W_v_up.weight.view(self.n_heads, self.head_dim, self.d_latent) o = torch.einsum('bthL,hdl->bthd', o_lat, Wv_up_T) # (B,T,32,128) return self.W_o(o.reshape(B,T,-1)), (c_K, c_V) # 缓存仍是 latent
v·k^T 把历史信息累加进一个固定大小的状态矩阵 S,新的历史不断写进去、旧的历史按遗忘门淡出。这样无论 L 多大,内存恒定 O(1),每步算力也恒定 O(1),不再随上下文增长。
k、v 算出来后立即写进状态 S(外积累加),Q 则作为读指针从 S 中取回压缩后的历史。Q、K、V 都在,只是“存储形态”从列表变成了状态机。
class KDA(nn.Module): # 对应 图KDA-1:完全替换图B蓝色框(线性注意力) def __init__(self, d_model=4096, s_dim=2048): self.s_dim = s_dim self.W_q = nn.Linear(d_model, s_dim) # 图KDA-2:q —— 读指针 self.W_k = nn.Linear(d_model, s_dim) # 图KDA-2:k —— 写地址 self.W_v = nn.Linear(d_model, s_dim) # 图KDA-2:v —— 写内容 self.W_beta = nn.Linear(d_model, s_dim) # 图KDA-2:δ —— 遗忘门(0~1) self.W_read = nn.Linear(s_dim, d_model) # 图KDA-2:读状态 → 输出 def forward(self, x, state=None): B = x.shape[0] # —— QKV 在哪里:Q=读、K=写地址、V=写内容,都在状态机内部 —— q = self.W_q(x).squeeze(1) # (B, s) 读取用 k = self.W_k(x).squeeze(1) # (B, s) 写地址 v = self.W_v(x).squeeze(1) # (B, s) 写内容 decay = torch.sigmoid(self.W_beta(x)).squeeze(1) # (B, s) 遗忘率 0~1 if state is None: state = torch.zeros(B, self.s_dim, self.s_dim, device=x.device) # 图KDA-1 状态更新:S ← δ·S + v·k^T(外积累加,历史信息写进固定状态) state = decay.unsqueeze(1) * state + torch.einsum('bs,bt->bst', v, k) # (B,s,s) # 图KDA-1 读取:out = W_read(q·S),从状态读回压缩后的全部历史 out = self.W_read(torch.einsum('bs,bst->bt', q, state)).unsqueeze(1) # (B,1,d_model) return out, state # 状态大小恒定,与 L 无关 → O(1) 显存、O(1) 算力
| 机制 | 每 token KV | @4096 总缓存 | 相对 MHA | 每步注意力 | 省内存靠 | 省算力靠 | 质量 |
|---|---|---|---|---|---|---|---|
| MHA | 16KB | 64MB | 100%(基线) | O(L) ≈33.6M | — | — | ⭐⭐⭐⭐⭐ |
| MQA | 0.5KB | 2MB | -97% | O(L) ≈33.6M | 个数→1 | K/V投影 ÷32 | ⭐⭐⭐ |
| GQA | 2KB | 8MB | -87.5% | O(L) ≈33.6M | 个数→G=4 | K/V投影 ÷8 + 带宽 ÷8 | ⭐⭐⭐⭐ |
| MLA | 1KB | 4MB | -93.75% | O(L) 但省展开 | 尺寸→d_latent | K/V投影 ÷16 + 解压吸收 | ⭐⭐⭐⭐⭐ |
| KDA | 恒定 8MB | 8MB(恒定) | 与 L 无关 | O(1) | 不存 KV 列表 | 线性注意力 ÷L | ⭐⭐⭐⭐ |