"""模型主体:Transformer Block、MTP 多 token 预测头、完整模型。""" import math from typing import Optional, Tuple, List import torch import torch.nn as nn import torch.nn.functional as F from .layers import RMSNorm, build_rope_cache from .mla import MLA from .moe import MoE, DenseFFN def _shift_left(x: torch.Tensor, k: int, fill: int): """y[:, i] = x[:, i+k],越界位置填 fill。""" if k == 0: return x.clamp_min(0) if fill == 0 else x y = torch.full_like(x, fill) if k < x.size(1): y[:, :-k] = x[:, k:] if fill == 0: y = y.clamp_min(0) return y class Block(nn.Module): """Pre-Norm 残差块:x + MLA(norm(x)) → x + FFN(norm(x))。""" def __init__(self, cfg, layer_id: int): super().__init__() self.attn_norm = RMSNorm(cfg.dim, cfg.norm_eps) self.attn = MLA(cfg) self.ffn_norm = RMSNorm(cfg.dim, cfg.norm_eps) # V3 的做法:靠前的层用稠密 FFN,后面的层才换成 MoE,训练早期更稳 self.is_moe = layer_id >= cfg.n_dense_layers self.ffn = MoE(cfg) if self.is_moe else DenseFFN(cfg) def forward(self, x, cos, sin, start_pos=0, mask=None): x = x + self.attn(self.attn_norm(x), cos, sin, start_pos, mask) h, aux = self.ffn(self.ffn_norm(x)) return x + h, aux class MTPModule(nn.Module): """Multi-Token Prediction 模块(DeepSeek-V3)。 普通语言模型每个位置只预测下一个 token,训练信号比较稀疏。 MTP 再挂一个轻量模块,让位置 i 同时预测 i+2: h'_i = Block( W · [ RMSNorm(h_i) ; RMSNorm(Emb(t_{i+1})) ] ) p(t_{i+2}) = Head(RMSNorm(h'_i)) 嵌入层和输出头与主干共享,所以额外开销只有一层 Block。 好处有二:训练时数据效率更高(模型被迫规划更远); 推理时这个头天然就是个 draft model,可以做自投机解码,实测约 1.8x 加速。 """ def __init__(self, cfg, layer_id: int): super().__init__() self.h_norm = RMSNorm(cfg.dim, cfg.norm_eps) self.e_norm = RMSNorm(cfg.dim, cfg.norm_eps) self.proj = nn.Linear(2 * cfg.dim, cfg.dim, bias=False) self.block = Block(cfg, layer_id) def forward(self, h, emb, cos, sin, start_pos=0, mask=None): z = self.proj(torch.cat([self.h_norm(h), self.e_norm(emb)], dim=-1)) return self.block(z, cos, sin, start_pos, mask) class DeepSeekFlashMini(nn.Module): def __init__(self, cfg): super().__init__() self.cfg = cfg self.embed = nn.Embedding(cfg.vocab_size, cfg.dim) self.layers = nn.ModuleList([Block(cfg, i) for i in range(cfg.n_layers)]) self.norm = RMSNorm(cfg.dim, cfg.norm_eps) self.head = nn.Linear(cfg.dim, cfg.vocab_size, bias=False) if cfg.tie_embeddings: self.head.weight = self.embed.weight self.mtp = nn.ModuleList([MTPModule(cfg, cfg.n_layers + i) for i in range(cfg.n_mtp)]) self.mtp_norm = RMSNorm(cfg.dim, cfg.norm_eps) if cfg.n_mtp > 0 else None cos, sin = build_rope_cache(cfg.qk_rope_head_dim, cfg.max_seq_len, cfg.rope_theta, cfg.rope_scaling) self.register_buffer("rope_cos", cos, persistent=False) self.register_buffer("rope_sin", sin, persistent=False) self.apply(self._init_weights) # 残差分支的输出投影按深度缩放初始化,防止深层激活方差爆掉 for name, p in self.named_parameters(): if name.endswith("wo.weight") or name.endswith("w_down.weight"): nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * cfg.n_layers)) self._cache_ready = False @staticmethod def _init_weights(m): if isinstance(m, nn.Linear): nn.init.normal_(m.weight, mean=0.0, std=0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.Embedding): nn.init.normal_(m.weight, mean=0.0, std=0.02) # ================= 前向 ================= def _rope(self, start_pos: int, T: int): if start_pos + T > self.cfg.max_seq_len: raise ValueError( f"位置 {start_pos + T} 超出上下文上限 max_seq_len={self.cfg.max_seq_len}。" f"请加大配置里的 max_seq_len(并考虑设置 rope_scaling 做外推)。") return self.rope_cos[start_pos:start_pos + T], self.rope_sin[start_pos:start_pos + T] def _mask(self, T: int, start_pos: int, device): """带 cache 的增量解码时,为多 token 输入构造对齐的因果 mask。""" if T == 1: return None if not self._cache_ready: return None # 训练路径:交给 SDPA 的 is_causal S = start_pos + T return torch.ones(T, S, dtype=torch.bool, device=device).tril(S - T) def forward_trunk(self, tokens: torch.Tensor, start_pos: int = 0): """主干前向,返回 (最后一层隐状态 h, logits, 路由辅助损失)。""" B, T = tokens.shape cos, sin = self._rope(start_pos, T) mask = self._mask(T, start_pos, tokens.device) h = self.embed(tokens) aux_total = h.new_zeros(()) for layer in self.layers: h, aux = layer(h, cos, sin, start_pos, mask) aux_total = aux_total + aux logits = self.head(self.norm(h)) return h, logits, aux_total def mtp_forward(self, h: torch.Tensor, next_tokens: torch.Tensor, start_pos: int = 0, depth: int = 0): """MTP 第 depth 层:给定主干隐状态和「下一个 token」,预测再下一个。""" B, T = next_tokens.shape cos, sin = self._rope(start_pos, T) mask = self._mask(T, start_pos, next_tokens.device) emb = self.embed(next_tokens) h2, aux = self.mtp[depth](h, emb, cos, sin, start_pos, mask) logits = self.head(self.mtp_norm(h2)) return h2, logits, aux def forward(self, tokens: torch.Tensor, targets: Optional[torch.Tensor] = None, start_pos: int = 0): """训练入口。targets[i] 应为 tokens[i+1](由 dataset 准备好)。""" h, logits, aux = self.forward_trunk(tokens, start_pos) if targets is None: return {"logits": logits} loss_main = F.cross_entropy(logits.reshape(-1, logits.size(-1)), targets.reshape(-1), ignore_index=-100) out = {"logits": logits, "loss_main": loss_main, "aux_loss": aux} # ---- MTP:第 d 个模块用 t_{i+1+d} 当输入,监督目标是 t_{i+2+d} ---- loss_mtp = logits.new_zeros(()) if self.cfg.n_mtp > 0: cur_h = h for d in range(self.cfg.n_mtp): inp = _shift_left(targets, d, fill=0) # t_{i+1+d} tgt = _shift_left(targets, d + 1, fill=-100) # t_{i+2+d} cur_h, mtp_logits, aux_d = self.mtp_forward(cur_h, inp, start_pos, d) aux = aux + aux_d loss_mtp = loss_mtp + F.cross_entropy( mtp_logits.reshape(-1, mtp_logits.size(-1)), tgt.reshape(-1), ignore_index=-100) loss_mtp = loss_mtp / self.cfg.n_mtp out["loss_mtp"] = loss_mtp out["aux_loss"] = aux out["loss"] = loss_main + self.cfg.mtp_loss_weight * loss_mtp + aux return out # ================= KV cache 管理 ================= def setup_cache(self, max_batch: int, max_seq_len: int, device=None, dtype=None): device = device or self.embed.weight.device dtype = dtype or self.embed.weight.dtype for m in self.modules(): if isinstance(m, MLA): m.setup_cache(max_batch, max_seq_len, device, dtype) self._cache_ready = True def clear_cache(self): for m in self.modules(): if isinstance(m, MLA): m.clear_cache() self._cache_ready = False def set_attn_impl(self, impl: str): for m in self.modules(): if isinstance(m, MLA): m.attn_impl = impl # ================= 统计 ================= def num_params(self) -> dict: cfg = self.cfg total = sum(p.numel() for p in self.parameters()) if cfg.tie_embeddings: total_unique = total else: total_unique = total # 激活参数量:每个 token 实际参与计算的那部分 inactive = 0 for layer in self.layers: if layer.is_moe: per_expert = sum(p.numel() for p in layer.ffn.experts[0].parameters()) skipped = cfg.n_routed_experts - cfg.n_activated_experts inactive += per_expert * skipped mtp_params = sum(p.numel() for p in self.mtp.parameters()) return { "total": total_unique, "activated": total_unique - inactive - mtp_params, # 推理时 MTP 可不用 "mtp": mtp_params, "embedding": self.embed.weight.numel(), } def kv_cache_bytes_per_token(self, dtype_size: int = 2) -> dict: """对比 MLA 与等规模 MHA 的 KV cache 开销。""" cfg = self.cfg mla = (cfg.kv_lora_rank + cfg.qk_rope_head_dim) * cfg.n_layers * dtype_size mha = 2 * cfg.n_heads * cfg.v_head_dim * cfg.n_layers * dtype_size return {"mla": mla, "mha_equivalent": mha, "ratio": mha / mla} def expert_load_stats(self) -> List[float]: return [layer.ffn.gate.load_stats()["imbalance"] for layer in self.layers if layer.is_moe] @torch.no_grad() def update_expert_bias(self): """训练循环每个 optimizer step 后调用一次。""" speed = self.cfg.bias_update_speed if speed <= 0: return for m in self.modules(): if isinstance(m, MoE): m.gate.update_bias(speed)