File size: 10,357 Bytes
7c97e31
d640b79
 
 
 
 
 
 
 
 
 
 
 
 
7c97e31
 
 
 
e85601b
 
 
d640b79
 
e85601b
 
 
d640b79
e85601b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d640b79
e85601b
d640b79
e85601b
 
 
7c97e31
e85601b
 
 
7c97e31
 
d640b79
 
7c97e31
d640b79
 
 
 
 
e85601b
 
d640b79
 
e85601b
d640b79
e85601b
d640b79
 
 
 
 
 
 
e85601b
d640b79
 
 
e85601b
 
 
d640b79
 
 
 
 
 
 
 
e85601b
 
 
d640b79
 
 
 
 
 
e85601b
d640b79
e85601b
 
d640b79
 
e85601b
d640b79
 
 
 
e85601b
 
 
 
d640b79
 
e85601b
 
d640b79
e85601b
 
 
d640b79
e85601b
d640b79
e85601b
 
d640b79
e85601b
d640b79
e85601b
 
 
 
 
 
 
d640b79
e85601b
d640b79
e85601b
 
d640b79
 
 
e85601b
 
 
 
 
 
 
 
d640b79
 
 
 
e85601b
d640b79
 
e85601b
7c97e31
 
e85601b
7c97e31
 
e85601b
 
 
 
 
d640b79
 
 
e85601b
 
d640b79
e85601b
 
 
d640b79
e85601b
d640b79
 
 
 
 
 
 
 
e85601b
 
 
 
 
d640b79
e85601b
d640b79
 
e85601b
 
7c97e31
 
e85601b
7c97e31
d640b79
7c97e31
 
e85601b
7c97e31
e85601b
 
 
7c97e31
 
e85601b
7c97e31
e85601b
 
7c97e31
 
e85601b
7c97e31
 
d640b79
7c97e31
 
 
 
 
d640b79
7c97e31
 
d640b79
7c97e31
e85601b
 
 
7c97e31
 
 
e85601b
7c97e31
 
 
 
 
e85601b
7c97e31
 
d640b79
e85601b
 
 
d640b79
 
 
 
e85601b
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
"""
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)  # (seq_len, head_dim//2)
    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)  # (1, 1, T, head_dim)
    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  # 每个 KV head 服务多少个 Q head

        # 分开三组投影:Q 全宽,K/V 收窄到 n_kv_head * head_dim
        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)

        # QK-norm:每个 head 独立归一化,稳定大 batch / 长训练时的 attention logits
        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)

        # QK-norm(在 RoPE 之前)
        q = self.q_norm(q)
        k = self.k_norm(k)

        q, k = apply_rope(q, k, cos[:T], sin[:T])

        # GQA:把 K/V 在 head 维度上重复 n_rep 次以匹配 Q
        # 比 SDPA 的 enable_gqa 参数更通用(兼容旧版 PyTorch)
        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)

        # 权重绑定:lm_head 与 token_emb 共享,节省 ~26M 参数
        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)
        # 残差流出口:attention 的 o_proj、MLP 的 c_proj
        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()

    # 去除 torch.compile 添加的 '_orig_mod.' 前缀
    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("接口检查通过。")