simak31 commited on
Commit
f4a7ec3
·
verified ·
1 Parent(s): 4b6dbd0

Upload model.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. model.py +173 -0
model.py ADDED
@@ -0,0 +1,173 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Same architecture family as before (RMSNorm, RoPE, grouped-query attention
3
+ via F.scaled_dot_product_attention, SwiGLU, tied embeddings) -- only the
4
+ CONFIG changed (see configs/config.py): small custom vocab, shorter context,
5
+ sized to land ~18.9M params at a genuine ~20:1 token:param ratio.
6
+
7
+ Run directly to print exact param count + smoke test:
8
+ python model.py
9
+ """
10
+ import torch
11
+ import torch.nn as nn
12
+ import torch.nn.functional as F
13
+
14
+ from configs.config import ModelConfig
15
+
16
+
17
+ class RMSNorm(nn.Module):
18
+ def __init__(self, dim: int, eps: float = 1e-5):
19
+ super().__init__()
20
+ self.eps = eps
21
+ self.weight = nn.Parameter(torch.ones(dim))
22
+
23
+ def forward(self, x):
24
+ norm = x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
25
+ return norm * self.weight
26
+
27
+
28
+ def precompute_rope(head_dim, seq_len, theta, device, dtype=torch.float32):
29
+ freqs = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device, dtype=dtype) / head_dim))
30
+ t = torch.arange(seq_len, device=device, dtype=dtype)
31
+ freqs = torch.outer(t, freqs)
32
+ return torch.cos(freqs), torch.sin(freqs)
33
+
34
+
35
+ def apply_rope(x, cos, sin):
36
+ x1, x2 = x[..., 0::2], x[..., 1::2]
37
+ cos = cos[None, None, :, :]
38
+ sin = sin[None, None, :, :]
39
+ r1 = x1 * cos - x2 * sin
40
+ r2 = x1 * sin + x2 * cos
41
+ return torch.stack([r1, r2], dim=-1).flatten(-2).to(x.dtype)
42
+
43
+
44
+ class GQAttention(nn.Module):
45
+ def __init__(self, cfg: ModelConfig):
46
+ super().__init__()
47
+ assert cfg.d_model % cfg.n_head == 0
48
+ assert cfg.n_head % cfg.n_kv_head == 0
49
+ self.n_head = cfg.n_head
50
+ self.n_kv_head = cfg.n_kv_head
51
+ self.head_dim = cfg.d_model // cfg.n_head
52
+ self.n_rep = cfg.n_head // cfg.n_kv_head
53
+
54
+ self.q_proj = nn.Linear(cfg.d_model, cfg.n_head * self.head_dim, bias=False)
55
+ self.k_proj = nn.Linear(cfg.d_model, cfg.n_kv_head * self.head_dim, bias=False)
56
+ self.v_proj = nn.Linear(cfg.d_model, cfg.n_kv_head * self.head_dim, bias=False)
57
+ self.o_proj = nn.Linear(cfg.n_head * self.head_dim, cfg.d_model, bias=False)
58
+
59
+ def forward(self, x, cos, sin):
60
+ b, t, _ = x.shape
61
+ q = self.q_proj(x).view(b, t, self.n_head, self.head_dim).transpose(1, 2)
62
+ k = self.k_proj(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
63
+ v = self.v_proj(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
64
+ q = apply_rope(q, cos, sin)
65
+ k = apply_rope(k, cos, sin)
66
+ if self.n_rep > 1:
67
+ k = k.repeat_interleave(self.n_rep, dim=1)
68
+ v = v.repeat_interleave(self.n_rep, dim=1)
69
+ out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
70
+ out = out.transpose(1, 2).contiguous().view(b, t, self.n_head * self.head_dim)
71
+ return self.o_proj(out)
72
+
73
+
74
+ class SwiGLU(nn.Module):
75
+ def __init__(self, cfg: ModelConfig):
76
+ super().__init__()
77
+ self.gate_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
78
+ self.up_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
79
+ self.down_proj = nn.Linear(cfg.d_ff, cfg.d_model, bias=False)
80
+
81
+ def forward(self, x):
82
+ return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
83
+
84
+
85
+ class Block(nn.Module):
86
+ def __init__(self, cfg: ModelConfig):
87
+ super().__init__()
88
+ self.attn_norm = RMSNorm(cfg.d_model)
89
+ self.attn = GQAttention(cfg)
90
+ self.mlp_norm = RMSNorm(cfg.d_model)
91
+ self.mlp = SwiGLU(cfg)
92
+ self.dropout = nn.Dropout(cfg.dropout)
93
+
94
+ def forward(self, x, cos, sin):
95
+ x = x + self.dropout(self.attn(self.attn_norm(x), cos, sin))
96
+ x = x + self.dropout(self.mlp(self.mlp_norm(x)))
97
+ return x
98
+
99
+
100
+ class TinyTransformer(nn.Module):
101
+ def __init__(self, cfg: ModelConfig):
102
+ super().__init__()
103
+ self.cfg = cfg
104
+ self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
105
+ self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layer)])
106
+ self.final_norm = RMSNorm(cfg.d_model)
107
+ self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
108
+ if cfg.tie_embeddings:
109
+ self.lm_head.weight = self.tok_emb.weight
110
+ self.head_dim = cfg.d_model // cfg.n_head
111
+ self.apply(self._init_weights)
112
+
113
+ def _init_weights(self, module):
114
+ if isinstance(module, nn.Linear):
115
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
116
+ if module.bias is not None:
117
+ nn.init.zeros_(module.bias)
118
+ elif isinstance(module, nn.Embedding):
119
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
120
+
121
+ def forward(self, idx, targets=None):
122
+ b, t = idx.shape
123
+ assert t <= self.cfg.context_len, f"seq len {t} exceeds context_len {self.cfg.context_len}"
124
+ cos, sin = precompute_rope(self.head_dim, t, self.cfg.rope_theta, idx.device)
125
+ cos, sin = cos.to(self.tok_emb.weight.dtype), sin.to(self.tok_emb.weight.dtype)
126
+ x = self.tok_emb(idx)
127
+ for block in self.blocks:
128
+ x = block(x, cos, sin)
129
+ x = self.final_norm(x)
130
+ logits = self.lm_head(x)
131
+ loss = None
132
+ if targets is not None:
133
+ loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=-1)
134
+ return logits, loss
135
+
136
+ @torch.no_grad()
137
+ def generate(self, idx, max_new_tokens, temperature=1.0, top_k=None):
138
+ for _ in range(max_new_tokens):
139
+ idx_cond = idx if idx.size(1) <= self.cfg.context_len else idx[:, -self.cfg.context_len:]
140
+ logits, _ = self(idx_cond)
141
+ logits = logits[:, -1, :] / max(temperature, 1e-5)
142
+ if top_k is not None:
143
+ v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
144
+ logits[logits < v[:, [-1]]] = -float("inf")
145
+ probs = F.softmax(logits, dim=-1)
146
+ next_id = torch.multinomial(probs, num_samples=1)
147
+ idx = torch.cat([idx, next_id], dim=1)
148
+ return idx
149
+
150
+ def num_params(self, non_embedding=False):
151
+ n = sum(p.numel() for p in self.parameters())
152
+ if non_embedding:
153
+ n -= self.tok_emb.weight.numel()
154
+ return n
155
+
156
+
157
+ if __name__ == "__main__":
158
+ cfg = ModelConfig()
159
+ model = TinyTransformer(cfg)
160
+ n = model.num_params()
161
+ n_emb = model.tok_emb.weight.numel()
162
+ print(f"Config: vocab={cfg.vocab_size} d_model={cfg.d_model} n_layer={cfg.n_layer} "
163
+ f"n_head={cfg.n_head} n_kv_head={cfg.n_kv_head} d_ff={cfg.d_ff} context_len={cfg.context_len}")
164
+ print(f"Total parameters: {n:,} (~{n/1e6:.2f}M)")
165
+ print(f"Embedding: {n_emb:,} ({100*n_emb/n:.0f}% of total)")
166
+
167
+ x = torch.randint(0, cfg.vocab_size, (2, 64))
168
+ y = torch.randint(0, cfg.vocab_size, (2, 64))
169
+ logits, loss = model(x, y)
170
+ assert logits.shape == (2, 64, cfg.vocab_size)
171
+ loss.backward()
172
+ n_missing = sum(1 for p in model.parameters() if p.grad is None)
173
+ print(f"Forward/backward OK. loss={loss.item():.3f} params_without_grad={n_missing}")