musabc commited on
Commit
04d9347
·
verified ·
1 Parent(s): 073320f

upload 05_train_v5.py

Browse files
Files changed (1) hide show
  1. 05_train_v5.py +594 -0
05_train_v5.py ADDED
@@ -0,0 +1,594 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ V5 Eğitim — 200M model, multi-stage curriculum (SmolLM2 tarzı).
3
+ Hedef GPU: RTX PRO 6000 Blackwell (96GB VRAM).
4
+
5
+ Mimari (model_v5.py):
6
+ - 18 layer, 14 head, 896 embd, 32K vocab, 1024 context
7
+ - RoPE + RMSNorm + SwiGLU + QK-norm + soft-cap + tied embeddings
8
+
9
+ Stack:
10
+ - Muon (2D weights) + AdamW (1D + embedding)
11
+ - bf16 + TF32 + cudnn.benchmark + torch.compile
12
+ - Async prefetcher per stage
13
+ - Multi-stage data loader (weighted sampling, progresif)
14
+
15
+ Curriculum (SmolLM2 stil) — Stage1=web(3-4B), Stage2=medium(9B), Stage3=premium(3B):
16
+ Faz 1 [0% - 55%] : Stage1 %25, Stage2 %65, Stage3 %10 (medium bulk + web)
17
+ Faz 2 [55% - 85%] : Stage1 %15, Stage2 %55, Stage3 %30 (premium ısınma)
18
+ Faz 3 [85% - 100%] : Stage1 %5 , Stage2 %25, Stage3 %70 (PREMIUM annealing)
19
+
20
+ Replay: Geçmiş stage'leri tamamen kesmiyoruz → catastrophic forgetting önlenir.
21
+
22
+ Kullanim:
23
+ python 05_train_v5.py # bastan
24
+ python 05_train_v5.py --resume # latest_ckpt
25
+ python 05_train_v5.py --compile # torch.compile
26
+ python 05_train_v5.py --max-time 480 # 8 saatlik oturum
27
+
28
+ Önceden: data/v5_stage1.bin, v5_stage2.bin, v5_stage3.bin, v5_val.bin hazır olmalı.
29
+ """
30
+
31
+ import argparse
32
+ import math
33
+ import os
34
+ import signal
35
+ import sys
36
+ import time
37
+ from contextlib import nullcontext
38
+ from pathlib import Path
39
+
40
+ import numpy as np
41
+ import torch
42
+ from tokenizers import Tokenizer
43
+
44
+ from model_v5 import GPTV5, GPTConfigV5
45
+ from muon import Muon
46
+
47
+ # ============================================================
48
+ # Konfigurasyon — V5 (RTX PRO 6000 Blackwell 96GB için)
49
+ # ============================================================
50
+ DATA_DIR = Path(__file__).parent / "data"
51
+ OUT_DIR = Path(__file__).parent / "runs" / "tr-200m-v5"
52
+
53
+ MODEL_CONFIG = dict(
54
+ block_size=2048,
55
+ vocab_size=32000,
56
+ n_layer=18,
57
+ n_head=14,
58
+ n_embd=896,
59
+ dropout=0.0,
60
+ rope_theta=10000.0,
61
+ logit_softcap=30.0,
62
+ )
63
+
64
+ # RTX PRO 6000 Blackwell 96GB VRAM
65
+ # T=2048 için activation memory T=1024'ün ~2x'i (Flash-Attn ile O(B*H*T))
66
+ # bs=48, T=2048 → activations ~60-70GB, weights+opt ~5GB → ~70GB toplam (rahat)
67
+ # 48 = 16*3 tensor core dostu
68
+ BATCH_SIZE = 32
69
+ GRAD_ACCUM_STEPS = 16 # etkin batch = 528, token/step ≈ 1.08M
70
+ MAX_STEPS = 20_000 # ~21.6B token training (Modern overtraining, 108x Chinchilla)
71
+ LOG_INTERVAL = 10
72
+ EVAL_INTERVAL = 400
73
+ EVAL_ITERS = 80
74
+ SAVE_INTERVAL = 1000
75
+ SAMPLE_INTERVAL = 2000
76
+
77
+ # LR — 200M, etkin batch ~528 için
78
+ MUON_LR = 0.022
79
+ ADAM_LR = 3.5e-4
80
+ MIN_LR_RATIO = 0.1
81
+ WARMUP_STEPS = 1000 # 20K step için %5 warmup
82
+ LR_DECAY_STEPS = 20_000
83
+
84
+ # Optimizer
85
+ WEIGHT_DECAY = 0.1 # 200M model için biraz weight decay yararlı
86
+ ADAM_BETA1 = 0.9
87
+ ADAM_BETA2 = 0.95
88
+ MUON_MOMENTUM = 0.95
89
+ GRAD_CLIP = 1.0
90
+
91
+ # Curriculum faz sınırları (oran cinsinden)
92
+ # Stage1 = WEB (oscar, mc4, forum, fineweb_hq) ~3-4B token
93
+ # Stage2 = MEDIUM (bellaturca, cosmos, culturax, havadis, cosmopedia) ~9B token (BULK)
94
+ # Stage3 = PREMIUM (wiki, wikisource, tezler, akademik, finepdfs, ozenli) ~3B token
95
+ PHASE1_END = 0.55 # 0-55% : bulk learning (medium dominant)
96
+ PHASE2_END = 0.85 # 55-85% : premium ısınır
97
+ # 85-100% : PREMIUM annealing
98
+
99
+ # Faz başına karışım oranları [stage1=web, stage2=medium, stage3=premium]
100
+ PHASE_MIX = {
101
+ 1: (0.25, 0.65, 0.10), # medium bulk + web, premium dokun
102
+ 2: (0.15, 0.55, 0.30), # premium ısın
103
+ 3: (0.05, 0.25, 0.70), # PREMIUM annealing — son finishing
104
+ }
105
+ # ============================================================
106
+
107
+ LATEST_CKPT = OUT_DIR / "latest_ckpt.pt"
108
+ BEST_CKPT = OUT_DIR / "best_ckpt.pt"
109
+ LOG_FILE = OUT_DIR / "train.log"
110
+
111
+
112
+ def get_lr_factor(step: int) -> float:
113
+ if step < WARMUP_STEPS:
114
+ return (step + 1) / (WARMUP_STEPS + 1)
115
+ if step > LR_DECAY_STEPS:
116
+ return MIN_LR_RATIO
117
+ decay_ratio = (step - WARMUP_STEPS) / (LR_DECAY_STEPS - WARMUP_STEPS)
118
+ coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))
119
+ return MIN_LR_RATIO + coeff * (1.0 - MIN_LR_RATIO)
120
+
121
+
122
+ def get_phase(step: int) -> int:
123
+ """Mevcut step'e göre curriculum faz (1/2/3)."""
124
+ p = step / max(MAX_STEPS, 1)
125
+ if p < PHASE1_END:
126
+ return 1
127
+ if p < PHASE2_END:
128
+ return 2
129
+ return 3
130
+
131
+
132
+ def log(msg: str):
133
+ print(msg, flush=True)
134
+ try:
135
+ with open(LOG_FILE, "a", encoding="utf-8") as f:
136
+ f.write(msg + "\n")
137
+ except Exception:
138
+ pass
139
+
140
+
141
+ # =====================================================================
142
+ # Data — per-stage loader
143
+ # =====================================================================
144
+ class StageLoader:
145
+ def __init__(self, bin_path: Path, block_size: int, batch_size: int,
146
+ device: torch.device, pin: bool = True, name: str = ""):
147
+ self.data = np.memmap(bin_path, dtype=np.uint16, mode="r")
148
+ self.block_size = block_size
149
+ self.batch_size = batch_size
150
+ self.device = device
151
+ self.pin = pin and device.type == "cuda"
152
+ self.name = name or bin_path.stem
153
+ n_tok = len(self.data)
154
+ log(f" {bin_path.name}: {n_tok:,} token (~{n_tok*2/1e9:.2f} GB)")
155
+ self.n_tokens = n_tok
156
+
157
+ def get_batch(self):
158
+ bs, T = self.batch_size, self.block_size
159
+ ix = np.random.randint(0, len(self.data) - T - 1, size=bs)
160
+ x_np = np.empty((bs, T), dtype=np.int64)
161
+ y_np = np.empty((bs, T), dtype=np.int64)
162
+ for k, i in enumerate(ix):
163
+ x_np[k] = self.data[i:i+T]
164
+ y_np[k] = self.data[i+1:i+1+T]
165
+ x = torch.from_numpy(x_np)
166
+ y = torch.from_numpy(y_np)
167
+ if self.pin:
168
+ x = x.pin_memory()
169
+ y = y.pin_memory()
170
+ x = x.to(self.device, non_blocking=True)
171
+ y = y.to(self.device, non_blocking=True)
172
+ return x, y
173
+
174
+
175
+ class MultiStageLoader:
176
+ """Curriculum-aware sampler — fazlara göre stage karışımı değişir."""
177
+ def __init__(self, stage_loaders, rng=None):
178
+ # stage_loaders: [s1, s2, s3]
179
+ self.loaders = stage_loaders
180
+ self.rng = rng or np.random.default_rng()
181
+
182
+ def get_batch(self, phase: int):
183
+ mix = PHASE_MIX[phase]
184
+ # Tek bir stage seçimi (batch içi karışım değil — daha temiz gradient)
185
+ idx = self.rng.choice(len(self.loaders), p=mix)
186
+ return self.loaders[idx].get_batch(), idx
187
+
188
+
189
+ class AsyncMultiStagePrefetcher:
190
+ """Faz bilgisi geçilen prefetch kuyruğu. Her get() çağrısında mevcut phase
191
+ kullanılır — geriden gelen batch'ler hâlâ önceki phase'in karışımındaysa
192
+ sorun değil (geçişler yumuşaktır)."""
193
+ def __init__(self, multi_loader: MultiStageLoader, phase_fn, queue_size=4):
194
+ import threading, queue
195
+ self.ml = multi_loader
196
+ self.phase_fn = phase_fn
197
+ self.q = queue.Queue(maxsize=queue_size)
198
+ self._stop = threading.Event()
199
+ self.thread = threading.Thread(target=self._produce, daemon=True)
200
+ self.thread.start()
201
+
202
+ def _produce(self):
203
+ while not self._stop.is_set():
204
+ try:
205
+ ph = self.phase_fn()
206
+ self.q.put(self.ml.get_batch(ph))
207
+ except Exception:
208
+ self._stop.set()
209
+ break
210
+
211
+ def get_batch(self):
212
+ return self.q.get()
213
+
214
+ def close(self):
215
+ self._stop.set()
216
+
217
+
218
+ # =====================================================================
219
+ # Eval / Sample
220
+ # =====================================================================
221
+ @torch.no_grad()
222
+ def estimate_loss(model, val_loader, train_loaders, ctx, eval_iters: int):
223
+ """Val + her stage için train loss."""
224
+ out = {}
225
+ model.eval()
226
+
227
+ # Val
228
+ losses = torch.zeros(eval_iters)
229
+ for k in range(eval_iters):
230
+ x, y = val_loader.get_batch()
231
+ with ctx:
232
+ _, loss = model(x, y)
233
+ losses[k] = loss.item()
234
+ out["val"] = losses.mean().item()
235
+
236
+ # Her stage'den birkaç iter
237
+ n_small = max(eval_iters // 4, 8)
238
+ for i, ld in enumerate(train_loaders, start=1):
239
+ losses = torch.zeros(n_small)
240
+ for k in range(n_small):
241
+ x, y = ld.get_batch()
242
+ with ctx:
243
+ _, loss = model(x, y)
244
+ losses[k] = loss.item()
245
+ out[f"stage{i}"] = losses.mean().item()
246
+
247
+ model.train()
248
+ return out
249
+
250
+
251
+ @torch.no_grad()
252
+ def sample_text(model, tokenizer, device, ctx,
253
+ prompt: str = "Türkiye", max_new_tokens: int = 100,
254
+ temperature: float = 0.8, top_k: int = 50,
255
+ repetition_penalty: float = 1.15):
256
+ model.eval()
257
+ ids = tokenizer.encode(prompt).ids
258
+ x = torch.tensor([ids], dtype=torch.long, device=device)
259
+ real_model = model._orig_mod if hasattr(model, "_orig_mod") else model
260
+ with ctx:
261
+ out = real_model.generate(
262
+ x, max_new_tokens=max_new_tokens,
263
+ temperature=temperature, top_k=top_k,
264
+ repetition_penalty=repetition_penalty,
265
+ )
266
+ text = tokenizer.decode(out[0].tolist())
267
+ model.train()
268
+ return text
269
+
270
+
271
+ # =====================================================================
272
+ # Checkpointing
273
+ # =====================================================================
274
+ def atomic_save(state: dict, path: Path):
275
+ tmp = path.with_suffix(path.suffix + ".tmp")
276
+ torch.save(state, tmp)
277
+ if path.exists():
278
+ path.unlink()
279
+ tmp.rename(path)
280
+
281
+
282
+ def build_state(model, opt_muon, opt_adam, scaler, step, best_val):
283
+ real_model = model._orig_mod if hasattr(model, "_orig_mod") else model
284
+ return {
285
+ "model": real_model.state_dict(),
286
+ "opt_muon": opt_muon.state_dict(),
287
+ "opt_adam": opt_adam.state_dict(),
288
+ "scaler": scaler.state_dict(),
289
+ "step": step,
290
+ "best_val": best_val,
291
+ "config": MODEL_CONFIG,
292
+ "version": "v5",
293
+ }
294
+
295
+
296
+ # =====================================================================
297
+ # Optimizer setup
298
+ # =====================================================================
299
+ def create_optimizers(model, device):
300
+ muon_params = []
301
+ adam_params = []
302
+
303
+ for name, p in model.named_parameters():
304
+ if not p.requires_grad:
305
+ continue
306
+ if p.ndim < 2:
307
+ adam_params.append(p)
308
+ elif "wte" in name or "lm_head" in name:
309
+ adam_params.append(p)
310
+ else:
311
+ muon_params.append(p)
312
+
313
+ seen = set()
314
+ adam_params_unique = []
315
+ for p in adam_params:
316
+ if id(p) not in seen:
317
+ seen.add(id(p))
318
+ adam_params_unique.append(p)
319
+
320
+ n_muon = sum(p.numel() for p in muon_params)
321
+ n_adam = sum(p.numel() for p in adam_params_unique)
322
+ log(f" Muon params: {n_muon/1e6:.2f}M ({len(muon_params)} tensor)")
323
+ log(f" AdamW params: {n_adam/1e6:.2f}M ({len(adam_params_unique)} tensor)")
324
+
325
+ opt_muon = Muon(
326
+ muon_params,
327
+ lr=MUON_LR,
328
+ momentum=MUON_MOMENTUM,
329
+ nesterov=True,
330
+ ns_steps=5,
331
+ )
332
+ opt_adam = torch.optim.AdamW(
333
+ adam_params_unique,
334
+ lr=ADAM_LR,
335
+ betas=(ADAM_BETA1, ADAM_BETA2),
336
+ weight_decay=WEIGHT_DECAY,
337
+ fused=(device.type == "cuda"),
338
+ )
339
+ return opt_muon, opt_adam
340
+
341
+
342
+ # =====================================================================
343
+ # Main
344
+ # =====================================================================
345
+ def main():
346
+ parser = argparse.ArgumentParser()
347
+ parser.add_argument("--resume", action="store_true")
348
+ parser.add_argument("--resume-best", action="store_true")
349
+ parser.add_argument("--compile", action="store_true")
350
+ parser.add_argument("--max-time", type=int, default=0,
351
+ help="Maksimum süre (dakika)")
352
+ parser.add_argument("--max-steps", type=int, default=None)
353
+ parser.add_argument("--batch", type=int, default=None,
354
+ help="Override BATCH_SIZE")
355
+ parser.add_argument("--grad-accum", type=int, default=None)
356
+ args = parser.parse_args()
357
+
358
+ OUT_DIR.mkdir(parents=True, exist_ok=True)
359
+
360
+ global MAX_STEPS, LR_DECAY_STEPS, BATCH_SIZE, GRAD_ACCUM_STEPS
361
+ if args.max_steps:
362
+ MAX_STEPS = args.max_steps
363
+ LR_DECAY_STEPS = args.max_steps
364
+ if args.batch:
365
+ BATCH_SIZE = args.batch
366
+ if args.grad_accum:
367
+ GRAD_ACCUM_STEPS = args.grad_accum
368
+
369
+ # Cihaz
370
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
371
+ if device.type == "cpu":
372
+ log("UYARI: CUDA yok.")
373
+ else:
374
+ log(f"GPU: {torch.cuda.get_device_name(0)}")
375
+ log(f"CUDA: {torch.version.cuda}, PyTorch: {torch.__version__}")
376
+ torch.set_float32_matmul_precision("high")
377
+ torch.backends.cuda.matmul.allow_tf32 = True
378
+ torch.backends.cudnn.allow_tf32 = True
379
+ torch.backends.cudnn.benchmark = True
380
+ # Blackwell: Flash Attention v2/v3 backend zorla (sdpa içinden)
381
+ try:
382
+ torch.backends.cuda.enable_flash_sdp(True)
383
+ torch.backends.cuda.enable_mem_efficient_sdp(True)
384
+ torch.backends.cuda.enable_math_sdp(False)
385
+ except Exception:
386
+ pass
387
+ # Daha agresif allocator (büyük bs için fragmentasyon azalır)
388
+ os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF",
389
+ "expandable_segments:True")
390
+ log("Perf: TF32 ON, cudnn.benchmark ON, Flash-SDPA ON")
391
+ # VRAM rapor
392
+ vram_gb = torch.cuda.get_device_properties(0).total_memory / 1e9
393
+ log(f"VRAM: {vram_gb:.1f} GB")
394
+
395
+ use_bf16 = device.type == "cuda" and torch.cuda.is_bf16_supported()
396
+ dtype = torch.bfloat16 if use_bf16 else torch.float16
397
+ log(f"Mixed precision: {dtype}")
398
+ ctx = (nullcontext() if device.type == "cpu"
399
+ else torch.amp.autocast(device_type="cuda", dtype=dtype))
400
+
401
+ # Data — 3 stage + val
402
+ log("\nData yukleniyor...")
403
+ bs, T = BATCH_SIZE, MODEL_CONFIG["block_size"]
404
+ stage1 = StageLoader(DATA_DIR / "v5_stage1.bin", T, bs, device, name="stage1")
405
+ stage2 = StageLoader(DATA_DIR / "v5_stage2.bin", T, bs, device, name="stage2")
406
+ stage3 = StageLoader(DATA_DIR / "v5_stage3.bin", T, bs, device, name="stage3")
407
+ val_loader = StageLoader(DATA_DIR / "v5_val.bin", T, bs, device, name="val")
408
+
409
+ total_tokens = stage1.n_tokens + stage2.n_tokens + stage3.n_tokens
410
+ log(f" Toplam: {total_tokens/1e9:.2f}B token")
411
+
412
+ multi = MultiStageLoader([stage1, stage2, stage3])
413
+
414
+ # Prefetcher — phase fonksiyonunu bir mutable ref ile vereceğiz
415
+ step_ref = {"step": 0}
416
+ def cur_phase():
417
+ return get_phase(step_ref["step"])
418
+
419
+ prefetch = AsyncMultiStagePrefetcher(multi, cur_phase, queue_size=4)
420
+ log(" Async multi-stage prefetcher aktif")
421
+
422
+ tokenizer = Tokenizer.from_file(str(DATA_DIR / "tokenizer-tr-v5.json"))
423
+
424
+ # Model
425
+ log("\nModel V5 olusturuluyor...")
426
+ cfg = GPTConfigV5(**MODEL_CONFIG)
427
+ model = GPTV5(cfg).to(device)
428
+ n_params = model.num_params()
429
+ log(f" Toplam: {n_params/1e6:.2f}M param")
430
+ log(f" Mimari: RoPE + RMSNorm + SwiGLU + QK-norm + soft-cap + tied emb")
431
+ log(f" L={cfg.n_layer}, H={cfg.n_head}, d={cfg.n_embd}, T={cfg.block_size}")
432
+
433
+ # Optimizers
434
+ opt_muon, opt_adam = create_optimizers(model, device)
435
+ log(f" Muon LR: {MUON_LR}, Momentum: {MUON_MOMENTUM}")
436
+ log(f" AdamW LR: {ADAM_LR}, WD: {WEIGHT_DECAY}")
437
+
438
+ scaler = torch.amp.GradScaler("cuda", enabled=(dtype == torch.float16))
439
+
440
+ # Resume
441
+ start_step = 0
442
+ best_val = float("inf")
443
+ resume_path = None
444
+ if args.resume_best and BEST_CKPT.exists():
445
+ resume_path = BEST_CKPT
446
+ elif args.resume and LATEST_CKPT.exists():
447
+ resume_path = LATEST_CKPT
448
+
449
+ if resume_path:
450
+ log(f"\nResume: {resume_path}")
451
+ ckpt = torch.load(resume_path, map_location=device, weights_only=False)
452
+ if ckpt.get("version") != "v5":
453
+ log("UYARI: V5 olmayan checkpoint!")
454
+ model.load_state_dict(ckpt["model"])
455
+ opt_muon.load_state_dict(ckpt["opt_muon"])
456
+ opt_adam.load_state_dict(ckpt["opt_adam"])
457
+ if "scaler" in ckpt:
458
+ scaler.load_state_dict(ckpt["scaler"])
459
+ start_step = ckpt["step"] + 1
460
+ best_val = ckpt.get("best_val", float("inf"))
461
+ log(f" step={start_step}, best_val={best_val:.4f}")
462
+
463
+ # Compile
464
+ if args.compile:
465
+ log("torch.compile baslatiliyor...")
466
+ torch._dynamo.config.suppress_errors = False
467
+ os.environ.setdefault("TORCHINDUCTOR_CACHE_DIR",
468
+ str(OUT_DIR / "_inductor_cache"))
469
+ model = torch.compile(model, mode="default", dynamic=False)
470
+
471
+ # Sinyal
472
+ interrupt_flag = {"stop": False}
473
+ def signal_handler(sig, frame):
474
+ if interrupt_flag["stop"]:
475
+ log("\n[!] Ikinci Ctrl+C, cikiyor.")
476
+ sys.exit(1)
477
+ interrupt_flag["stop"] = True
478
+ log("\n[!] Ctrl+C alindi, kaydedilip cikilacak.")
479
+ signal.signal(signal.SIGINT, signal_handler)
480
+
481
+ total_tokens_per_step = BATCH_SIZE * GRAD_ACCUM_STEPS * MODEL_CONFIG["block_size"]
482
+ log(f"\nEgitim basliyor:")
483
+ log(f" Step araligi: {start_step} → {MAX_STEPS}")
484
+ log(f" Etkin batch: {BATCH_SIZE * GRAD_ACCUM_STEPS}")
485
+ log(f" Token/step: {total_tokens_per_step:,}")
486
+ log(f" Toplam token: {MAX_STEPS * total_tokens_per_step / 1e9:.1f}B")
487
+ log(f" Curriculum: P1[0-{int(PHASE1_END*100)}%] "
488
+ f"P2[{int(PHASE1_END*100)}-{int(PHASE2_END*100)}%] "
489
+ f"P3[{int(PHASE2_END*100)}-100%]")
490
+
491
+ t_start = time.time()
492
+ step_t0 = time.time()
493
+ step = start_step
494
+ last_phase = -1
495
+ stage_hits = [0, 0, 0]
496
+
497
+ try:
498
+ while step < MAX_STEPS:
499
+ step_ref["step"] = step
500
+ phase = get_phase(step)
501
+ if phase != last_phase:
502
+ mix = PHASE_MIX[phase]
503
+ log(f"\n>>> FAZ {phase} basliyor (step {step}): "
504
+ f"stage1={mix[0]:.0%}, stage2={mix[1]:.0%}, stage3={mix[2]:.0%}")
505
+ last_phase = phase
506
+
507
+ # LR
508
+ lr_factor = get_lr_factor(step)
509
+ muon_lr = MUON_LR * lr_factor
510
+ adam_lr = ADAM_LR * lr_factor
511
+ for pg in opt_muon.param_groups:
512
+ pg["lr"] = muon_lr
513
+ for pg in opt_adam.param_groups:
514
+ pg["lr"] = adam_lr
515
+
516
+ # Grad accumulation
517
+ opt_muon.zero_grad(set_to_none=True)
518
+ opt_adam.zero_grad(set_to_none=True)
519
+ loss_accum = 0.0
520
+ for _ in range(GRAD_ACCUM_STEPS):
521
+ (x, y), stage_idx = prefetch.get_batch()
522
+ stage_hits[stage_idx] += 1
523
+ with ctx:
524
+ _, loss = model(x, y)
525
+ loss = loss / GRAD_ACCUM_STEPS
526
+ scaler.scale(loss).backward()
527
+ loss_accum += loss.item()
528
+
529
+ scaler.unscale_(opt_adam)
530
+ torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)
531
+ opt_muon.step()
532
+ scaler.step(opt_adam)
533
+ scaler.update()
534
+
535
+ # Log
536
+ if step % LOG_INTERVAL == 0:
537
+ dt = time.time() - step_t0
538
+ tps = (LOG_INTERVAL * total_tokens_per_step) / dt if step > start_step else 0
539
+ step_t0 = time.time()
540
+ elapsed_min = (time.time() - t_start) / 60
541
+ total_hits = sum(stage_hits) or 1
542
+ mix_str = "/".join(f"{h*100//total_hits}" for h in stage_hits)
543
+ log(f"step {step:>6} | P{phase} | loss {loss_accum:.4f} | "
544
+ f"muon {muon_lr:.2e} adam {adam_lr:.2e} | "
545
+ f"{tps/1e3:.0f}K tok/s | mix {mix_str} | {elapsed_min:.1f}m")
546
+ stage_hits = [0, 0, 0]
547
+
548
+ # Eval
549
+ if step > start_step and step % EVAL_INTERVAL == 0:
550
+ losses = estimate_loss(model, val_loader,
551
+ [stage1, stage2, stage3], ctx, EVAL_ITERS)
552
+ log(f" >>> EVAL: val {losses['val']:.4f} "
553
+ f"s1 {losses['stage1']:.4f} s2 {losses['stage2']:.4f} "
554
+ f"s3 {losses['stage3']:.4f}")
555
+ if losses["val"] < best_val:
556
+ best_val = losses["val"]
557
+ state = build_state(model, opt_muon, opt_adam, scaler, step, best_val)
558
+ atomic_save(state, BEST_CKPT)
559
+ log(f" >>> BEST kaydedildi (val {best_val:.4f})")
560
+
561
+ # Save
562
+ if step > start_step and step % SAVE_INTERVAL == 0:
563
+ state = build_state(model, opt_muon, opt_adam, scaler, step, best_val)
564
+ atomic_save(state, LATEST_CKPT)
565
+
566
+ # Sample
567
+ if step > start_step and step % SAMPLE_INTERVAL == 0:
568
+ for prompt in ["Türkiye", "Yapay zeka", "Bu çalışmada"]:
569
+ text = sample_text(model, tokenizer, device, ctx,
570
+ prompt=prompt, max_new_tokens=80)
571
+ log(f" [sample] {text!r}")
572
+
573
+ # Time
574
+ if args.max_time and (time.time() - t_start) / 60 >= args.max_time:
575
+ log(f"\n[time] {args.max_time} dakika doldu, kaydedilip cikiliyor.")
576
+ break
577
+
578
+ if interrupt_flag["stop"]:
579
+ break
580
+
581
+ step += 1
582
+
583
+ finally:
584
+ log("\nSon checkpoint yaziliyor...")
585
+ state = build_state(model, opt_muon, opt_adam, scaler, step, best_val)
586
+ atomic_save(state, LATEST_CKPT)
587
+ log(f" latest_ckpt.pt → step {step}, best_val {best_val:.4f}")
588
+ prefetch.close()
589
+
590
+ log(f"\n[DONE] Step {step}/{MAX_STEPS}. Best val: {best_val:.4f}")
591
+
592
+
593
+ if __name__ == "__main__":
594
+ main()