Compactbot commited on
Commit
0686dc8
·
verified ·
1 Parent(s): 8aca2b5

Add ram-18m training script (18,290,304 params, LLaMA-style GQA+SwiGLU)

Browse files
Files changed (1) hide show
  1. train_ram_18m.py +431 -0
train_ram_18m.py ADDED
@@ -0,0 +1,431 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """ram-18m: 18.3M-param LLaMA-style language model trained from scratch.
3
+
4
+ Architecture:
5
+ d_model=384, n_heads=6, n_kv_heads=2 (GQA), n_layers=7
6
+ SwiGLU FFN (4x), RoPE, RMSNorm, vocab 8192, tied embed/head, ctx 512
7
+ Total: 18,290,304 learnable parameters
8
+
9
+ Default training:
10
+ ~2B tokens (FineWeb-Edu L3), AdamW 2e-4, cosine + warmup
11
+ batch 32 (effective), seq 512, ~12,207 steps
12
+
13
+ Usage:
14
+ python3 train_ram_18m.py --stage prepare # download + tokenize data
15
+ python3 train_ram_18m.py --stage train # train the model
16
+ python3 train_ram_18m.py --stage all
17
+ python3 train_ram_18m.py --stage eval --ckpt path/to/ckpt.pt
18
+
19
+ Requirements:
20
+ pip install torch transformers datasets numpy tokenizers
21
+ """
22
+ import os, sys, math, json, time, argparse, glob, random
23
+ import numpy as np
24
+ import torch
25
+ import torch.nn as nn
26
+ import torch.nn.functional as F
27
+ from torch.optim import AdamW
28
+
29
+ # ============================================================================
30
+ # Architecture
31
+ # ============================================================================
32
+ VOCAB = 8192
33
+ D_MODEL = 384
34
+ N_HEADS = 6
35
+ N_KV_HEADS = 2
36
+ N_LAYERS = 7
37
+ HEAD_DIM = D_MODEL // N_HEADS # 64
38
+ KV_DIM = N_KV_HEADS * HEAD_DIM # 128
39
+ FFN_DIM = D_MODEL * 4 # 1536
40
+ SEQ_LEN = 512
41
+ ROPE_THETA = 10000.0
42
+
43
+ # Verified param count: 18,290,304
44
+
45
+
46
+ class RMSNorm(nn.Module):
47
+ def __init__(self, dim, eps=1e-6):
48
+ super().__init__()
49
+ self.eps = eps
50
+ self.weight = nn.Parameter(torch.ones(dim))
51
+
52
+ def forward(self, x):
53
+ norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
54
+ return (x.float() * norm).type_as(x) * self.weight
55
+
56
+
57
+ class RoPE(nn.Module):
58
+ def __init__(self, head_dim, theta=ROPE_THETA):
59
+ super().__init__()
60
+ freqs = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim))
61
+ self.register_buffer("freqs", freqs, persistent=False)
62
+
63
+ def forward(self, x, pos):
64
+ freqs = self.freqs
65
+ angles = pos[:, None].float() * freqs[None, :]
66
+ cos = angles.cos()[None, None, :, :]
67
+ sin = angles.sin()[None, None, :, :]
68
+ x1 = x[..., 0::2]
69
+ x2 = x[..., 1::2]
70
+ out1 = x1 * cos - x2 * sin
71
+ out2 = x1 * sin + x2 * cos
72
+ return torch.stack([out1, out2], dim=-1).flatten(-2)
73
+
74
+
75
+ class GQAAttention(nn.Module):
76
+ def __init__(self):
77
+ super().__init__()
78
+ self.q_proj = nn.Linear(D_MODEL, N_HEADS * HEAD_DIM, bias=False)
79
+ self.k_proj = nn.Linear(D_MODEL, N_KV_HEADS * HEAD_DIM, bias=False)
80
+ self.v_proj = nn.Linear(D_MODEL, N_KV_HEADS * HEAD_DIM, bias=False)
81
+ self.o_proj = nn.Linear(N_HEADS * HEAD_DIM, D_MODEL, bias=False)
82
+ self.rope = RoPE(HEAD_DIM)
83
+
84
+ def forward(self, x, mask=None):
85
+ B, T, _ = x.shape
86
+ q = self.q_proj(x).view(B, T, N_HEADS, HEAD_DIM).transpose(1, 2)
87
+ k = self.k_proj(x).view(B, T, N_KV_HEADS, HEAD_DIM).transpose(1, 2)
88
+ v = self.v_proj(x).view(B, T, N_KV_HEADS, HEAD_DIM).transpose(1, 2)
89
+ pos = torch.arange(T, device=x.device)
90
+ q = self.rope(q, pos)
91
+ k = self.rope(k, pos)
92
+ rep = N_HEADS // N_KV_HEADS
93
+ k = k.repeat_interleave(rep, dim=1)
94
+ v = v.repeat_interleave(rep, dim=1)
95
+ scale = HEAD_DIM ** -0.5
96
+ attn = (q @ k.transpose(-2, -1)) * scale
97
+ if mask is not None:
98
+ attn = attn.masked_fill(mask[:, None, None, :] == 0, float("-inf"))
99
+ attn = F.softmax(attn, dim=-1)
100
+ out = attn @ v
101
+ out = out.transpose(1, 2).contiguous().view(B, T, N_HEADS * HEAD_DIM)
102
+ return self.o_proj(out)
103
+
104
+
105
+ class SwiGLU(nn.Module):
106
+ def __init__(self):
107
+ super().__init__()
108
+ self.gate = nn.Linear(D_MODEL, FFN_DIM, bias=False)
109
+ self.up = nn.Linear(D_MODEL, FFN_DIM, bias=False)
110
+ self.down = nn.Linear(FFN_DIM, D_MODEL, bias=False)
111
+
112
+ def forward(self, x):
113
+ return self.down(F.silu(self.gate(x)) * self.up(x))
114
+
115
+
116
+ class TransformerBlock(nn.Module):
117
+ def __init__(self):
118
+ super().__init__()
119
+ self.attn_norm = RMSNorm(D_MODEL)
120
+ self.attn = GQAAttention()
121
+ self.ffn_norm = RMSNorm(D_MODEL)
122
+ self.ffn = SwiGLU()
123
+
124
+ def forward(self, x, mask=None):
125
+ x = x + self.attn(self.attn_norm(x), mask)
126
+ x = x + self.ffn(self.ffn_norm(x))
127
+ return x
128
+
129
+
130
+ class RAM18M(nn.Module):
131
+ def __init__(self):
132
+ super().__init__()
133
+ self.tok_emb = nn.Embedding(VOCAB, D_MODEL)
134
+ self.layers = nn.ModuleList([TransformerBlock() for _ in range(N_LAYERS)])
135
+ self.norm = RMSNorm(D_MODEL)
136
+ self.head = nn.Linear(D_MODEL, VOCAB, bias=False)
137
+ self.head.weight = self.tok_emb.weight
138
+ self.apply(self._init_weights)
139
+
140
+ def _init_weights(self, module):
141
+ if isinstance(module, nn.Linear):
142
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
143
+ if module.bias is not None:
144
+ nn.init.zeros_(module.bias)
145
+ elif isinstance(module, nn.Embedding):
146
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
147
+
148
+ def forward(self, input_ids, targets=None):
149
+ B, T = input_ids.shape
150
+ h = self.tok_emb(input_ids)
151
+ mask = torch.tril(torch.ones(T, T, device=input_ids.device))
152
+ for layer in self.layers:
153
+ h = layer(h, mask)
154
+ h = self.norm(h)
155
+ logits = self.head(h)
156
+ loss = None
157
+ if targets is not None:
158
+ loss = F.cross_entropy(logits.view(-1, VOCAB), targets.view(-1))
159
+ return logits, loss
160
+
161
+ def count_params(self):
162
+ return sum(p.numel() for p in self.parameters() if p.requires_grad)
163
+
164
+
165
+ def get_lr(step, max_steps, warmup, base_lr, min_lr):
166
+ if step < warmup:
167
+ return base_lr * (step + 1) / warmup
168
+ if step >= max_steps:
169
+ return min_lr
170
+ progress = (step - warmup) / (max_steps - warmup)
171
+ return min_lr + 0.5 * (base_lr - min_lr) * (1 + math.cos(math.pi * progress))
172
+
173
+
174
+ # ============================================================================
175
+ # Data
176
+ # ============================================================================
177
+ DATASET = "HuggingFaceFW/fineweb-edu"
178
+ DATASET_CONFIG = "sample-100BT"
179
+ SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
180
+ TOK_DIR = os.path.join(SCRIPT_DIR, "tokens")
181
+ CKPT_DIR = SCRIPT_DIR
182
+
183
+
184
+ def prepare_data(target_tokens=2_000_000_000):
185
+ """Download and tokenize FineWeb-Edu. Saves .npy files of token ids."""
186
+ from datasets import load_dataset
187
+ from tokenizers import Tokenizer
188
+ from tokenizers.models import BPE
189
+ from tokenizers.pre_tokenizers import Whitespace
190
+ from tokenizers.trainers import BpeTrainer
191
+
192
+ os.makedirs(TOK_DIR, exist_ok=True)
193
+
194
+ print("Training BPE tokenizer (vocab 8192)...")
195
+ ds = load_dataset(DATASET, DATASET_CONFIG, split="train", streaming=True)
196
+ texts = []
197
+ for i, row in enumerate(ds):
198
+ texts.append(row["text"])
199
+ if i >= 200000:
200
+ break
201
+
202
+ tokenizer = Tokenizer(BPE(unk_token="<unk>"))
203
+ tokenizer.pre_tokenizer = Whitespace()
204
+ trainer = BpeTrainer(vocab_size=VOCAB, special_tokens=["<unk>", "<pad>", "<bos>", "<eos>"])
205
+ tokenizer.train_from_iterator(texts, trainer)
206
+ tok_path = os.path.join(TOK_DIR, "tokenizer.json")
207
+ tokenizer.save(tok_path)
208
+ print(f"Tokenizer saved to {tok_path}")
209
+
210
+ print(f"Tokenizing up to {target_tokens:,} tokens...")
211
+ ds = load_dataset(DATASET, DATASET_CONFIG, split="train", streaming=True)
212
+ all_tokens = []
213
+ n_docs = 0
214
+ for row in ds:
215
+ ids = tokenizer.encode(row["text"])
216
+ if len(ids) < 10:
217
+ continue
218
+ all_tokens.extend(ids)
219
+ n_docs += 1
220
+ if len(all_tokens) >= target_tokens:
221
+ break
222
+ if n_docs % 100000 == 0:
223
+ print(f" {n_docs} docs, {len(all_tokens):,} tokens")
224
+
225
+ print(f"Total: {n_docs} docs, {len(all_tokens):,} tokens")
226
+ arr = np.array(all_tokens, dtype=np.int32)
227
+ part_size = 100_000_000
228
+ for i in range(0, len(arr), part_size):
229
+ part = arr[i:i+part_size]
230
+ path = os.path.join(TOK_DIR, f"part_{i//part_size:03d}.npy")
231
+ np.save(path, part)
232
+ print(f" Saved {path}: {len(part):,} tokens")
233
+ print("Data prep complete.")
234
+
235
+
236
+ class DataIterator:
237
+ """Streams tokenized data from .npy parts, yielding (input, target) batches."""
238
+ def __init__(self, tok_dir, batch_size, seq_len, device="cpu"):
239
+ self.parts = sorted(glob.glob(os.path.join(tok_dir, "part_*.npy")))
240
+ if not self.parts:
241
+ raise FileNotFoundError(f"No .npy files in {tok_dir}. Run --stage prepare first.")
242
+ self.batch_size = batch_size
243
+ self.seq_len = seq_len
244
+ self.device = device
245
+ self._buf = np.array([], dtype=np.int32)
246
+ self._part_idx = 0
247
+ self._rng = np.random.default_rng(42)
248
+
249
+ def _refill(self):
250
+ need = self.batch_size * self.seq_len + self.seq_len
251
+ while len(self._buf) < need:
252
+ if self._part_idx >= len(self.parts):
253
+ self._part_idx = 0
254
+ part = np.load(self.parts[self._part_idx], mmap_mode="r")
255
+ self._buf = np.concatenate([self._buf, np.array(part)])
256
+ self._part_idx += 1
257
+
258
+ def __iter__(self):
259
+ while True:
260
+ self._refill()
261
+ max_start = len(self._buf) - self.batch_size * self.seq_len - self.seq_len
262
+ if max_start < 0:
263
+ self._refill()
264
+ continue
265
+ start = int(self._rng.integers(0, max_start))
266
+ chunk = self._buf[start:start + self.batch_size * self.seq_len + self.seq_len]
267
+ flat = chunk.reshape(self.batch_size, self.seq_len + 1)
268
+ x = torch.tensor(flat[:, :-1], dtype=torch.long, device=self.device)
269
+ y = torch.tensor(flat[:, 1:], dtype=torch.long, device=self.device)
270
+ yield x, y
271
+
272
+
273
+ # ============================================================================
274
+ # Training
275
+ # ============================================================================
276
+ def train(
277
+ steps=12207,
278
+ batch_size=32,
279
+ seq_len=SEQ_LEN,
280
+ lr=2e-4,
281
+ min_lr=2e-5,
282
+ warmup=200,
283
+ accum=1,
284
+ ckpt_every=500,
285
+ eval_every=500,
286
+ device="cuda",
287
+ resume=None,
288
+ ):
289
+ """Train ram-18m from scratch."""
290
+ torch.manual_seed(42)
291
+ model = RAM18M()
292
+ n_params = model.count_params()
293
+ print(f"Model: {n_params:,} params")
294
+
295
+ start_step = 0
296
+ if resume:
297
+ ckpt = torch.load(resume, map_location="cpu")
298
+ model.load_state_dict(ckpt["model"])
299
+ start_step = ckpt["step"]
300
+ print(f"Resumed from {resume} at step {start_step}")
301
+
302
+ if device == "cuda" and torch.cuda.is_available():
303
+ model = model.cuda()
304
+ else:
305
+ device = "cpu"
306
+ model = model.to(device)
307
+ print(f"Device: {device}")
308
+
309
+ data_iter = DataIterator(TOK_DIR, batch_size, seq_len, device)
310
+ opt = AdamW(model.parameters(), lr=lr, betas=(0.9, 0.95), weight_decay=0.0)
311
+
312
+ model.train()
313
+ t0 = time.time()
314
+
315
+ for step in range(start_step, steps):
316
+ opt.zero_grad()
317
+ loss_accum = 0.0
318
+ for _ in range(accum):
319
+ x, y = next(iter(data_iter))
320
+ _, loss = model(x, y)
321
+ loss = loss / accum
322
+ loss.backward()
323
+ loss_accum += loss.item()
324
+
325
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
326
+ opt.step()
327
+ loss_val = loss_accum
328
+
329
+ if (step + 1) % 10 == 0:
330
+ elapsed = time.time() - t0
331
+ tok_per_sec = (step + 1 - start_step) * batch_size * seq_len / max(elapsed, 1)
332
+ lr_now = get_lr(step, steps, warmup, lr, min_lr)
333
+ print(f"step {step+1}/{steps} | loss {loss_val:.4f} | lr {lr_now:.6f} | {tok_per_sec:.0f} tok/s | {elapsed:.0f}s")
334
+
335
+ if (step + 1) % ckpt_every == 0:
336
+ path = os.path.join(CKPT_DIR, f"ckpt_step{step+1}.pt")
337
+ torch.save({"model": model.state_dict(), "opt": opt.state_dict(),
338
+ "step": step + 1, "loss": loss_val}, path)
339
+ print(f" checkpoint -> {path}")
340
+
341
+ if (step + 1) % eval_every == 0:
342
+ model.eval()
343
+ with torch.no_grad():
344
+ x, y = next(iter(data_iter))
345
+ _, eval_loss = model(x, y)
346
+ model.train()
347
+ print(f" eval_loss (1 batch): {eval_loss.item():.4f}")
348
+
349
+ path = os.path.join(CKPT_DIR, "final.pt")
350
+ torch.save({"model": model.state_dict(), "step": steps, "loss": loss_val}, path)
351
+ print(f"Training complete. Final model -> {path}")
352
+
353
+ # Print sample
354
+ model.eval()
355
+ torch.manual_seed(0)
356
+ with torch.no_grad():
357
+ prompt = torch.tensor([[3]], device=device)
358
+ for _ in range(200):
359
+ logits, _ = model(prompt)
360
+ next_tok = logits[0, -1].argmax()
361
+ prompt = torch.cat([prompt, next_tok.unsqueeze(0)], dim=1)
362
+ try:
363
+ from tokenizers import Tokenizer
364
+ tok_path = os.path.join(TOK_DIR, "tokenizer.json")
365
+ if os.path.exists(tok_path):
366
+ tok = Tokenizer.from_file(tok_path)
367
+ text = tok.decode(prompt[0].tolist())
368
+ print(f"\nSample generation:\n{text[:500]}")
369
+ except Exception:
370
+ pass
371
+
372
+
373
+ # ============================================================================
374
+ # Eval: zero-shot loglikelihood on standard benchmarks
375
+ # ============================================================================
376
+ def eval_benchmarks(ckpt_path, device="cuda", n_samples=500):
377
+ """Run zero-shot loglikelihood eval on PIQA, ARC-Easy, ARC-Challenge, HellaSwag."""
378
+ from datasets import load_dataset
379
+
380
+ model = RAM18M()
381
+ ckpt = torch.load(ckpt_path, map_location="cpu")
382
+ model.load_state_dict(ckpt["model"])
383
+ if device == "cuda" and torch.cuda.is_available():
384
+ model = model.cuda()
385
+ else:
386
+ device = "cpu"
387
+ model.eval()
388
+
389
+ tok_path = os.path.join(TOK_DIR, "tokenizer.json")
390
+ from tokenizers import Tokenizer
391
+ tokenizer = Tokenizer.from_file(tok_path)
392
+
393
+ def encode(text):
394
+ return tokenizer.encode(text).ids
395
+
396
+ def loglikelihood(context, continuation):
397
+ full_ids = encode(context + " " + continuation)
398
+ ctx_ids = encode(context)
399
+ ctx_len = min(len(ctx_ids), len(full_ids) - 1)
400
+ if ctx_len < 1:
401
+ return -1000.0
402
+ input_ids = torch.tensor([full_ids], device=device)
403
+ with torch.no_grad():
404
+ logits, _ = model(input_ids)
405
+ log_probs = F.log_softmax(logits[0, ctx_len-1:-1, :], dim=-1)
406
+ target_ids = torch.tensor(full_ids[ctx_len:], device=device)
407
+ if len(target_ids) == 0:
408
+ return -1000.0
409
+ return log_probs.gather(1, target_ids.unsqueeze(1)).sum().item()
410
+
411
+ def accuracy(items, n=n_samples):
412
+ correct = 0
413
+ total = 0
414
+ for item in items[:n]:
415
+ ctx = item["context"]
416
+ options = item["options"]
417
+ label = item["label"]
418
+ lls = [loglikelihood(ctx, opt) for opt in options]
419
+ pred = max(range(len(lls)), key=lambda i: lls[i])
420
+ if pred == label:
421
+ correct += 1
422
+ total += 1
423
+ if total % 50 == 0:
424
+ print(f" {total}/{n} done, acc so far: {100*correct/total:.1f}%")
425
+ return 100.0 * correct / max(total, 1)
426
+
427
+ results = {}
428
+
429
+ print("Loading PIQA...")
430
+ piqa = load_dataset("ybisk/piqa", split="validation")
431
+ piqa_items = [{"context":