| """ |
| SmolGPT — 参考 HuggingFace SmolLM2-135M 的 "深而窄" 设计 |
| |
| 设计理念(与同仓库 deep-gpt 对比): |
| - 深:n_layer 12 → 30(参考 SmolLM2 的 30 层) |
| - 窄:n_embd 768 → 512,head_dim 保持 64 |
| - GQA:分组查询注意力 8 Q heads / 2 KV heads(4:1) |
| 类比 SmolLM2 的 9 Q / 3 KV,减小 KV cache 占用 |
| - ReLU² 前馈(Primer 论文),hidden_dim=1728; |
| 相比 SwiGLU 少 1 个矩阵,同参数预算下更宽 |
| - RMSNorm + QK-norm,RoPE theta=100000(参考 SmolLM2) |
| - 残差出口零初始化(c_proj.weight = 0,起步即恒等映射)、tied embeddings |
| |
| 总参数:约 98.5M(满足 ≤100M 约束) |
| """ |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
|
|
| def precompute_freqs(head_dim: int, seq_len: int, theta: float = 100000.0): |
| """预计算 RoPE 旋转频率(cos/sin),形状 (seq_len, head_dim)。""" |
| inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)) |
| t = torch.arange(seq_len, dtype=torch.float32) |
| freqs = torch.outer(t, inv_freq) |
| cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1) |
| sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1) |
| return cos, sin |
|
|
|
|
| def rotate_half(x: torch.Tensor) -> torch.Tensor: |
| x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2 :] |
| return torch.cat((-x2, x1), dim=-1) |
|
|
|
|
| def apply_rope(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor): |
| cos = cos.unsqueeze(0).unsqueeze(0) |
| sin = sin.unsqueeze(0).unsqueeze(0) |
| q = q * cos + rotate_half(q) * sin |
| k = k * cos + rotate_half(k) * sin |
| return q, k |
|
|
|
|
| class RMSNorm(nn.Module): |
| """RMS 层归一化:相比 LayerNorm 无均值偏移,更快且数值稳定。""" |
|
|
| def __init__(self, dim: int, eps: float = 1e-5): |
| super().__init__() |
| self.eps = eps |
| self.weight = nn.Parameter(torch.ones(dim)) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt() |
| return (x.float() * norm * self.weight).type_as(x) |
|
|
|
|
| class GroupedQueryAttention(nn.Module): |
| """分组查询注意力(GQA)+ QK-norm + RoPE + Flash SDPA。 |
| |
| Q heads 数量 > KV heads 数量,K/V 在 head 维度上被多个 Q 共享。 |
| 显著降低注意力的参数量与推理 KV cache 占用。 |
| """ |
|
|
| def __init__(self, n_embd: int, n_head: int, n_kv_head: int): |
| super().__init__() |
| assert n_embd % n_head == 0, f"n_embd={n_embd} 必须能被 n_head={n_head} 整除" |
| assert n_head % n_kv_head == 0, f"n_head={n_head} 必须能被 n_kv_head={n_kv_head} 整除" |
|
|
| self.n_head = n_head |
| self.n_kv_head = n_kv_head |
| self.head_dim = n_embd // n_head |
| self.n_rep = n_head // n_kv_head |
|
|
| |
| self.q_proj = nn.Linear(n_embd, n_head * self.head_dim, bias=False) |
| self.k_proj = nn.Linear(n_embd, n_kv_head * self.head_dim, bias=False) |
| self.v_proj = nn.Linear(n_embd, n_kv_head * self.head_dim, bias=False) |
| self.o_proj = nn.Linear(n_embd, n_embd, bias=False) |
|
|
| |
| self.q_norm = RMSNorm(self.head_dim) |
| self.k_norm = RMSNorm(self.head_dim) |
|
|
| def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: |
| B, T, C = x.shape |
|
|
| q = self.q_proj(x).view(B, T, self.n_head, self.head_dim).transpose(1, 2) |
| k = self.k_proj(x).view(B, T, self.n_kv_head, self.head_dim).transpose(1, 2) |
| v = self.v_proj(x).view(B, T, self.n_kv_head, self.head_dim).transpose(1, 2) |
|
|
| |
| q = self.q_norm(q) |
| k = self.k_norm(k) |
|
|
| q, k = apply_rope(q, k, cos[:T], sin[:T]) |
|
|
| |
| |
| if self.n_rep > 1: |
| k = k.repeat_interleave(self.n_rep, dim=1) |
| v = v.repeat_interleave(self.n_rep, dim=1) |
|
|
| y = F.scaled_dot_product_attention(q, k, v, is_causal=True) |
| return self.o_proj(y.transpose(1, 2).contiguous().view(B, T, C)) |
|
|
|
|
| class ReLU2MLP(nn.Module): |
| """ReLU² 前馈:c_proj(ReLU(c_fc(x))²)。 |
| |
| 来自 Primer 论文 (https://arxiv.org/abs/2109.08668v2)。相比 SwiGLU: |
| - 只需 2 个矩阵(vs 3),同参数预算下 hidden_dim 可放大 1.5× |
| - 没有门控分支,前向/反向都更快 |
| - LM 任务上质量与 SwiGLU 持平或略优(modded-nanogpt 验证) |
| """ |
|
|
| def __init__(self, n_embd: int, hidden_dim: int): |
| super().__init__() |
| self.c_fc = nn.Linear(n_embd, hidden_dim, bias=False) |
| self.c_proj = nn.Linear(hidden_dim, n_embd, bias=False) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.c_proj(F.relu(self.c_fc(x)).square()) |
|
|
|
|
| class Block(nn.Module): |
| """Pre-norm 块:RMSNorm → GQA → 残差 → RMSNorm → ReLU² → 残差。""" |
|
|
| def __init__(self, n_embd: int, n_head: int, n_kv_head: int, hidden_dim: int): |
| super().__init__() |
| self.norm1 = RMSNorm(n_embd) |
| self.attn = GroupedQueryAttention(n_embd, n_head, n_kv_head) |
| self.norm2 = RMSNorm(n_embd) |
| self.ffn = ReLU2MLP(n_embd, hidden_dim) |
|
|
| def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: |
| x = x + self.attn(self.norm1(x), cos, sin) |
| x = x + self.ffn(self.norm2(x)) |
| return x |
|
|
|
|
| class SmolGPT(nn.Module): |
| """ |
| 深而窄的 GPT 模型,约 98.5M 参数。 |
| |
| 默认超参数: |
| n_embd=512, n_layer=30, n_head=8, n_kv_head=2 (head_dim=64), hidden_dim=1728 |
| (hidden_dim 从原 SwiGLU 的 1152 调到 ReLU² 的 1728,保持 MLP 总参数 |
| 不变:SwiGLU 用 3×512×1152 = ReLU² 用 2×512×1728) |
| """ |
|
|
| cos: torch.Tensor |
| sin: torch.Tensor |
|
|
| def __init__( |
| self, |
| vocab_size: int = 50257, |
| n_embd: int = 512, |
| n_head: int = 8, |
| n_kv_head: int = 2, |
| n_layer: int = 30, |
| block_size: int = 1024, |
| hidden_dim: int = 1728, |
| rope_theta: float = 100000.0, |
| ): |
| super().__init__() |
| self.block_size = block_size |
|
|
| self.token_emb = nn.Embedding(vocab_size, n_embd) |
|
|
| head_dim = n_embd // n_head |
| cos, sin = precompute_freqs(head_dim, block_size, theta=rope_theta) |
| self.register_buffer("cos", cos, persistent=False) |
| self.register_buffer("sin", sin, persistent=False) |
|
|
| self.blocks = nn.ModuleList( |
| [Block(n_embd, n_head, n_kv_head, hidden_dim) for _ in range(n_layer)] |
| ) |
| self.norm_f = RMSNorm(n_embd) |
|
|
| |
| self.lm_head = nn.Linear(n_embd, vocab_size, bias=False) |
| self.lm_head.weight = self.token_emb.weight |
|
|
| self._init_weights() |
|
|
| def _init_weights(self): |
| """GPT-2 风格初始化 + 残差出口零初始化。 |
| |
| 残差出口(GQA 的 o_proj、ReLU2MLP 的 c_proj)置零意味着训练起步时 |
| 每个 Block 都是恒等映射 `x ← x + 0`,模型等价于 Embed → Norm → LM_Head |
| 的浅层模型。残差贡献从第一步开始由训练过程精确学出,loss 曲线起步 |
| 更平滑(modded-nanogpt / Primer 等多个实验验证)。 |
| """ |
| for module in self.modules(): |
| if isinstance(module, nn.Linear): |
| nn.init.normal_(module.weight, mean=0.0, std=0.02) |
| elif isinstance(module, nn.Embedding): |
| nn.init.normal_(module.weight, mean=0.0, std=0.02) |
| |
| for name, p in self.named_parameters(): |
| if name.endswith(("o_proj.weight", "c_proj.weight")): |
| nn.init.zeros_(p) |
|
|
| def forward(self, input_ids: torch.Tensor) -> torch.Tensor: |
| """ |
| Args: |
| input_ids: LongTensor (B, T) |
| Returns: |
| logits: FloatTensor (B, T, vocab_size) |
| """ |
| B, T = input_ids.shape |
| assert T <= self.block_size, f"序列长度 {T} 超过 block_size {self.block_size}" |
|
|
| x = self.token_emb(input_ids) |
| cos = self.cos[:T] |
| sin = self.sin[:T] |
|
|
| for block in self.blocks: |
| x = block(x, cos, sin) |
|
|
| x = self.norm_f(x) |
| return self.lm_head(x) |
|
|
|
|
| |
|
|
| def load_model(checkpoint_path: str, device: str = "cuda") -> torch.nn.Module: |
| """从 checkpoint 加载 SmolGPT。支持纯 state_dict 或带 config 的字典格式。""" |
| checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True) |
|
|
| if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint: |
| state_dict = checkpoint["model_state_dict"] |
| config = checkpoint.get("config", {}) |
| model = SmolGPT(**config) |
| else: |
| state_dict = checkpoint |
| model = SmolGPT() |
|
|
| |
| unwanted_prefix = "_orig_mod." |
| for k in list(state_dict.keys()): |
| if k.startswith(unwanted_prefix): |
| state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k) |
|
|
| model.load_state_dict(state_dict) |
| model.to(device) |
| model.eval() |
| return model |
|
|
|
|
| |
|
|
| if __name__ == "__main__": |
| model = SmolGPT() |
|
|
| unique_params = sum(p.numel() for p in set(model.parameters())) |
| all_params = sum(p.numel() for p in model.parameters()) |
| print(f"唯一参数量: {unique_params:,} ({unique_params/1e6:.2f}M)") |
| print(f"全部参数量: {all_params:,} ({all_params/1e6:.2f}M) ← model.parameters() 统计") |
|
|
| assert unique_params <= 100_000_000, f"超过 100M 限制!实际 {unique_params/1e6:.2f}M" |
|
|
| dummy = torch.randint(0, 50257, (2, 1024)) |
| logits = model(dummy) |
| print(f"输入 shape: {dummy.shape}") |
| print(f"输出 shape: {logits.shape}") |
| assert logits.shape == (2, 1024, 50257), "输出形状错误!" |
| print("接口检查通过。") |
|
|