Compactbot commited on
Commit
66741e7
·
verified ·
1 Parent(s): 7c4a9a5

Add training script

Browse files
Files changed (1) hide show
  1. train_story10m.py +330 -0
train_story10m.py ADDED
@@ -0,0 +1,330 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ StoryLM-10M: ~10.5M param LLaMA-style language model trained on TinyStories.
4
+
5
+ Architecture:
6
+ d_model=256, n_heads=8, n_kv_heads=8 (MHA), n_layers=8
7
+ SwiGLU FFN (4x), RoPE, RMSNorm, vocab 8192, tied embed/head, ctx 512
8
+ Total: ~10.5M learnable parameters
9
+
10
+ Training:
11
+ ~2B tokens from TinyStories, AdamW 3e-4, cosine + warmup
12
+ batch 64 (effective), seq 512, ~6,100 steps
13
+
14
+ Usage:
15
+ python3 train_story10m.py --stage prepare # download + tokenize
16
+ python3 train_story10m.py --stage train # train
17
+ python3 train_story10m.py --stage eval # evaluate
18
+ python3 train_story10m.py --stage all
19
+ """
20
+ import os, sys, math, json, time, argparse, glob
21
+ import numpy as np
22
+ import torch
23
+ import torch.nn as nn
24
+ import torch.nn.functional as F
25
+
26
+ # ============================================================================
27
+ # Config
28
+ # ============================================================================
29
+ VOCAB = 8192
30
+ D_MODEL = 256
31
+ N_HEADS = 8
32
+ N_KV_HEADS = 8
33
+ N_LAYERS = 8
34
+ HEAD_DIM = D_MODEL // N_HEADS # 32
35
+ KV_DIM = N_KV_HEADS * HEAD_DIM # 256
36
+ FFN_DIM = D_MODEL * 4 # 1024
37
+ SEQ_LEN = 512
38
+ BATCH = 4
39
+ LR = 3e-4
40
+ WARMUP_STEPS = 500
41
+ TOTAL_STEPS = 24400
42
+ CHECKPOINT_EVERY = 500
43
+ OUT_DIR = "/work/story10m"
44
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
45
+ TOKENIZER_PATH = "/work/tokenizer.json"
46
+
47
+ os.makedirs(OUT_DIR, exist_ok=True)
48
+
49
+ # ============================================================================
50
+ # Architecture
51
+ # ============================================================================
52
+ class RMSNorm(nn.Module):
53
+ def __init__(self, dim, eps=1e-5):
54
+ super().__init__()
55
+ self.eps = eps
56
+ self.weight = nn.Parameter(torch.ones(dim))
57
+ def forward(self, x):
58
+ norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
59
+ return (x * norm).type_as(x) * self.weight
60
+
61
+ def precompute_rope(dim, max_seq, base=10000.0):
62
+ freqs = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
63
+ t = torch.arange(max_seq).float()
64
+ angles = torch.outer(t, freqs)
65
+ return torch.stack([angles.cos(), angles.sin()], dim=-1)
66
+
67
+ def apply_rope(x, cos, sin):
68
+ # x: (B, H, T, D)
69
+ B, H, T, D = x.shape
70
+ x1 = x[..., 0::2]
71
+ x2 = x[..., 1::2]
72
+ c = cos[:T].unsqueeze(0).unsqueeze(0) # (1, 1, T, D/2)
73
+ s = sin[:T].unsqueeze(0).unsqueeze(0)
74
+ out = torch.empty_like(x)
75
+ out[..., 0::2] = x1 * c - x2 * s
76
+ out[..., 1::2] = x1 * s + x2 * c
77
+ return out
78
+
79
+ class Attention(nn.Module):
80
+ def __init__(self, d_model, n_heads, n_kv_heads, head_dim):
81
+ super().__init__()
82
+ self.n_heads = n_heads
83
+ self.n_kv_heads = n_kv_heads
84
+ self.head_dim = head_dim
85
+ self.qkv = nn.Linear(d_model, (n_heads + 2 * n_kv_heads) * head_dim, bias=False)
86
+ self.o = nn.Linear(n_heads * head_dim, d_model, bias=False)
87
+
88
+ def forward(self, x, cos, sin):
89
+ B, T, _ = x.shape
90
+ qkv = self.qkv(x)
91
+ q, k, v = qkv.split([self.n_heads * self.head_dim,
92
+ self.n_kv_heads * self.head_dim,
93
+ self.n_kv_heads * self.head_dim], dim=-1)
94
+ q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
95
+ k = k.view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
96
+ v = v.view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
97
+ if self.n_kv_heads != self.n_heads:
98
+ rep = self.n_heads // self.n_kv_heads
99
+ k = k.repeat_interleave(rep, dim=1)
100
+ v = v.repeat_interleave(rep, dim=1)
101
+ q = apply_rope(q, cos, sin)
102
+ k = apply_rope(k, cos, sin)
103
+ attn = F.scaled_dot_product_attention(q, k, v, is_causal=True)
104
+ attn = attn.transpose(1, 2).contiguous().view(B, T, -1)
105
+ return self.o(attn)
106
+
107
+ class SwiGLU(nn.Module):
108
+ def __init__(self, d_model, ffn_dim):
109
+ super().__init__()
110
+ self.gate = nn.Linear(d_model, ffn_dim, bias=False)
111
+ self.up = nn.Linear(d_model, ffn_dim, bias=False)
112
+ self.down = nn.Linear(ffn_dim, d_model, bias=False)
113
+ def forward(self, x):
114
+ return self.down(F.silu(self.gate(x)) * self.up(x))
115
+
116
+ class Block(nn.Module):
117
+ def __init__(self, d_model, n_heads, n_kv_heads, head_dim, ffn_dim):
118
+ super().__init__()
119
+ self.norm1 = RMSNorm(d_model)
120
+ self.attn = Attention(d_model, n_heads, n_kv_heads, head_dim)
121
+ self.norm2 = RMSNorm(d_model)
122
+ self.ffn = SwiGLU(d_model, ffn_dim)
123
+ def forward(self, x, cos, sin):
124
+ x = x + self.attn(self.norm1(x), cos, sin)
125
+ x = x + self.ffn(self.norm2(x))
126
+ return x
127
+
128
+ class StoryLM(nn.Module):
129
+ def __init__(self):
130
+ super().__init__()
131
+ self.embed = nn.Embedding(VOCAB, D_MODEL)
132
+ self.layers = nn.ModuleList([
133
+ Block(D_MODEL, N_HEADS, N_KV_HEADS, HEAD_DIM, FFN_DIM)
134
+ for _ in range(N_LAYERS)
135
+ ])
136
+ self.norm = RMSNorm(D_MODEL)
137
+ self.lm_head = nn.Linear(D_MODEL, VOCAB, bias=False)
138
+ self.lm_head.weight = self.embed.weight # tied
139
+ rope = precompute_rope(HEAD_DIM, SEQ_LEN)
140
+ self.register_buffer('rope_cos', rope[..., 0])
141
+ self.register_buffer('rope_sin', rope[..., 1])
142
+
143
+ def forward(self, x, targets=None):
144
+ h = self.embed(x)
145
+ for block in self.layers:
146
+ h = block(h, self.rope_cos, self.rope_sin)
147
+ h = self.norm(h)
148
+ logits = self.lm_head(h)
149
+ loss = None
150
+ if targets is not None:
151
+ loss = F.cross_entropy(logits.view(-1, VOCAB), targets.view(-1))
152
+ return logits, loss
153
+
154
+ # ============================================================================
155
+ # Data
156
+ # ============================================================================
157
+ def prepare_data():
158
+ """Download TinyStories and tokenize."""
159
+ print("Loading tokenizer...")
160
+ from tokenizers import Tokenizer
161
+ tok = Tokenizer.from_file(TOKENIZER_PATH)
162
+
163
+ print("Downloading TinyStories...")
164
+ import urllib.request
165
+ stories_path = os.path.join(OUT_DIR, "tinystories.txt")
166
+ if not os.path.exists(stories_path):
167
+ url = "https://huggingface.co/datasets/roneneldan/TinyStories/resolve/main/TinyStoriesv2_2M.txt"
168
+ urllib.request.urlretrieve(url, stories_path)
169
+
170
+ with open(stories_path, 'r') as f:
171
+ stories = f.readlines()
172
+ print(f" {len(stories)} stories")
173
+
174
+ print("Tokenizing...")
175
+ all_tokens = []
176
+ for i, story in enumerate(stories):
177
+ story = story.strip()
178
+ if len(story) < 20:
179
+ continue
180
+ ids = tok.encode(story, add_special_tokens=False).ids
181
+ all_tokens.extend(ids)
182
+ if (i+1) % 100000 == 0:
183
+ print(f" {i+1}/{len(stories)} stories, {len(all_tokens)/1e6:.1f}M tokens so far")
184
+
185
+ tokens = np.array(all_tokens, dtype=np.uint16)
186
+ np.save(os.path.join(OUT_DIR, "tokens.npy"), tokens)
187
+ print(f" Total: {len(tokens)/1e6:.1f}M tokens")
188
+ print(f" Saved to {OUT_DIR}/tokens.npy")
189
+
190
+ def load_data():
191
+ tokens = np.load(os.path.join(OUT_DIR, "tokens.npy"))
192
+ print(f"Loaded {len(tokens)/1e6:.1f}M tokens")
193
+ # 99% train, 1% val
194
+ split = int(len(tokens) * 0.99)
195
+ train = tokens[:split]
196
+ val = tokens[split:]
197
+ print(f"Train: {len(train)/1e6:.1f}M, Val: {len(val)/1e6:.1f}M")
198
+ return train, val
199
+
200
+ def get_batch(data, batch_size, seq_len):
201
+ ix = torch.randint(len(data) - seq_len - 1, (batch_size,))
202
+ x = torch.stack([torch.from_numpy((data[i:i+seq_len]).astype(np.int64)) for i in ix])
203
+ y = torch.stack([torch.from_numpy((data[i+1:i+1+seq_len]).astype(np.int64)) for i in ix])
204
+ return x.to(DEVICE), y.to(DEVICE)
205
+
206
+ # ============================================================================
207
+ # Training
208
+ # ============================================================================
209
+ def train():
210
+ train_data, val_data = load_data()
211
+ model = StoryLM().to(DEVICE)
212
+ total_params = sum(p.numel() for p in model.parameters())
213
+ print(f"Parameters: {total_params:,}")
214
+
215
+ optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=0.0, betas=(0.9, 0.95))
216
+
217
+ best_val = float('inf')
218
+ t0 = time.time()
219
+
220
+ for step in range(1, TOTAL_STEPS + 1):
221
+ # LR schedule: warmup + cosine
222
+ if step < WARMUP_STEPS:
223
+ lr = LR * step / WARMUP_STEPS
224
+ else:
225
+ progress = (step - WARMUP_STEPS) / (TOTAL_STEPS - WARMUP_STEPS)
226
+ lr = LR * 0.5 * (1 + math.cos(math.pi * progress))
227
+ for param_group in optimizer.param_groups:
228
+ param_group['lr'] = lr
229
+
230
+ x, y = get_batch(train_data, BATCH, SEQ_LEN)
231
+ _, loss = model(x, y)
232
+ optimizer.zero_grad()
233
+ loss.backward()
234
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
235
+ optimizer.step()
236
+
237
+ if step % 100 == 0:
238
+ elapsed = time.time() - t0
239
+ tok_per_s = (step * BATCH * SEQ_LEN) / elapsed
240
+ print(f"step {step}/{TOTAL_STEPS} | loss {loss.item():.4f} | lr {lr:.2e} | {tok_per_s/1000:.1f}k tok/s | {elapsed:.0f}s", flush=True)
241
+
242
+ if step % CHECKPOINT_EVERY == 0:
243
+ # Eval
244
+ model.eval()
245
+ val_losses = []
246
+ for _ in range(20):
247
+ vx, vy = get_batch(val_data, BATCH, SEQ_LEN)
248
+ _, vloss = model(vx, vy)
249
+ val_losses.append(vloss.item())
250
+ val_loss = sum(val_losses) / len(val_losses)
251
+ model.train()
252
+ print(f" val_loss: {val_loss:.4f}")
253
+
254
+ ckpt_path = os.path.join(OUT_DIR, f"ckpt_{step}.pt")
255
+ torch.save(model.state_dict(), ckpt_path)
256
+ print(f" checkpoint: {ckpt_path}")
257
+
258
+ if val_loss < best_val:
259
+ best_val = val_loss
260
+ torch.save(model.state_dict(), os.path.join(OUT_DIR, "best.pt"))
261
+ print(f" new best: {val_loss:.4f}")
262
+
263
+ # Save final
264
+ torch.save(model.state_dict(), os.path.join(OUT_DIR, "final.pt"))
265
+ print(f"\nTraining done. Best val_loss: {best_val:.4f}")
266
+ print(f"Final model: {OUT_DIR}/final.pt")
267
+
268
+ # ============================================================================
269
+ # Evaluation
270
+ # ============================================================================
271
+ def evaluate():
272
+ model = StoryLM().to(DEVICE)
273
+ ckpt = torch.load(os.path.join(OUT_DIR, "best.pt"), map_location=DEVICE)
274
+ model.load_state_dict(ckpt)
275
+ model.eval()
276
+
277
+ _, val_data = load_data()
278
+
279
+ # Perplexity
280
+ val_losses = []
281
+ for _ in range(50):
282
+ vx, vy = get_batch(val_data, BATCH, SEQ_LEN)
283
+ _, vloss = model(vx, vy)
284
+ val_losses.append(vloss.item())
285
+ avg_loss = sum(val_losses) / len(val_losses)
286
+ ppl = math.exp(avg_loss)
287
+ print(f"Val perplexity: {ppl:.2f} (loss {avg_loss:.4f})")
288
+
289
+ # Generation samples
290
+ from tokenizers import Tokenizer
291
+ tok = Tokenizer.from_file(TOKENIZER_PATH)
292
+ prompts = [
293
+ "Once upon a time, there was a little",
294
+ "The cat sat on the",
295
+ "In the beginning, the world was",
296
+ "A small robot named",
297
+ "Every morning, the sun",
298
+ ]
299
+ print("\n=== Generation Samples ===")
300
+ for prompt in prompts:
301
+ ids = tok.encode(prompt, add_special_tokens=False).ids
302
+ x = torch.tensor([ids], dtype=torch.long, device=DEVICE)
303
+ with torch.no_grad():
304
+ for _ in range(30):
305
+ out, _ = model(x)
306
+ next_id = out[0, -1].argmax().item()
307
+ x = torch.cat([x, torch.tensor([[next_id]], device=DEVICE)], dim=1)
308
+ full_ids = ids + [next_id]
309
+ text = tok.decode(full_ids, skip_special_tokens=True)
310
+ print(f"Prompt: {prompt}")
311
+ print(f"Output: {text}")
312
+ print()
313
+
314
+ # ============================================================================
315
+ # Main
316
+ # ============================================================================
317
+ def main():
318
+ parser = argparse.ArgumentParser()
319
+ parser.add_argument("--stage", choices=["prepare", "train", "eval", "all"], default="all")
320
+ args = parser.parse_args()
321
+
322
+ if args.stage in ("prepare", "all"):
323
+ prepare_data()
324
+ if args.stage in ("train", "all"):
325
+ train()
326
+ if args.stage in ("eval", "all"):
327
+ evaluate()
328
+
329
+ if __name__ == "__main__":
330
+ main()