File size: 19,726 Bytes
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4e857fb
 
 
 
 
 
 
 
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1865e4d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1a1cdec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1a1cdec
 
 
 
 
 
 
 
4e857fb
 
 
 
 
1865e4d
 
 
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1a1cdec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4e857fb
 
 
 
 
 
 
 
 
 
 
 
0b47e1b
 
daec617
0b47e1b
 
 
 
 
daec617
 
 
 
 
 
 
 
 
 
 
 
 
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
daec617
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4e857fb
 
 
 
 
 
 
1a1cdec
 
 
 
 
 
 
daec617
 
 
 
 
0b47e1b
 
4e857fb
 
 
 
daec617
4e857fb
1a1cdec
 
 
 
 
 
 
daec617
 
0b47e1b
daec617
 
 
 
 
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1865e4d
 
 
 
 
 
 
 
0b47e1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
#!/usr/bin/env python
"""Fractus-1B cloud training script β€” GPU optimized, auto-push to HF.



Designed for RunPod/QuickPod RTX 4090 or A100.

Automatically uploads every checkpoint to HuggingFace.



Usage on cloud GPU:

    pip install torch datasets huggingface_hub

    python train_1b_cloud.py --epochs 20

    

Environment variables needed:

    HF_TOKEN=your_token_here

"""
import argparse, gc, math, os, sys, time
import torch, torch.nn as nn, torch.nn.functional as F
from fractus.model_1b import Fractus1B
from fractus.tokenizer import FractusTokenizer

# Try to import Triton kernels (will fail silently on CPU / no-triton).
try:
    from fractus.nn.triton_kernels import fused_linear_cross_entropy, TRITON_READY, self_test as triton_self_test
    _HAS_TRITON_IMPORT = True
except Exception:
    _HAS_TRITON_IMPORT = False
    TRITON_READY = False

HF_TOKEN = os.environ.get("HF_TOKEN", "")
HF_REPO = "thefinalboss/Fractus"


def upload_hf(path, repo_path):
    """Upload to HuggingFace. Never fails the training."""
    if not HF_TOKEN:
        print(f"  [HF] No token, skipping.", flush=True)
        return
    try:
        from huggingface_hub import HfApi
        api = HfApi(token=HF_TOKEN)
        api.upload_file(path_or_fileobj=path, path_in_repo=repo_path,
                       repo_id=HF_REPO, repo_type="model")
        print(f"  [HF] Uploaded {repo_path}", flush=True)
    except Exception as e:
        print(f"  [HF] Failed: {type(e).__name__} β€” training continues.", flush=True)


def save_and_upload(model, optimizer, epoch, loss, acc, config, ckpt_dir):
    """Save checkpoint locally + upload to HF."""
    os.makedirs(ckpt_dir, exist_ok=True)
    path = os.path.join(ckpt_dir, f"fractus_1b_epoch{epoch}.pt")
    torch.save({
        "model_state": model.state_dict(),
        "optimizer_state": optimizer.state_dict(),
        "config": config,
        "epoch": epoch,
        "loss": loss,
        "accuracy": acc,
        "timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
    }, path)
    size_mb = os.path.getsize(path) / 1e6
    print(f"  [ckpt] {path} ({size_mb:.0f}MB)", flush=True)
    upload_hf(path, f"checkpoints/fractus_1b_epoch{epoch}.pt")
    upload_hf(path, "checkpoints/fractus_1b_latest.pt")
    # Delete old checkpoint to save disk on cloud.
    if epoch > 1:
        old = os.path.join(ckpt_dir, f"fractus_1b_epoch{epoch-1}.pt")
        if os.path.exists(old):
            os.remove(old)
            print(f"  [disk] Removed old checkpoint {old}", flush=True)


def save_step_checkpoint(model, optimizer, step, epoch, loss, config, ckpt_dir,

                          keep_last=2):
    """Save a mid-epoch checkpoint every N steps + upload to HF.



    Named by global step (e.g. fractus_1b_step10000.pt). Uploads as both the

    step-named file AND 'fractus_1b_latest.pt' so resume always picks up the

    newest. Keeps only the last `keep_last` step checkpoints on disk to avoid

    filling the pod.



    This is the CRASH-RECOVERY path: if the pod dies mid-epoch, you resume

    from the latest step checkpoint and lose at most save_every steps of work.

    """
    os.makedirs(ckpt_dir, exist_ok=True)
    path = os.path.join(ckpt_dir, f"fractus_1b_step{step}.pt")
    torch.save({
        "model_state": model.state_dict(),
        "optimizer_state": optimizer.state_dict(),
        "config": config,
        "step": step,
        "epoch": epoch,
        "loss": loss,
        "timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
    }, path)
    size_mb = os.path.getsize(path) / 1e6
    print(f"  [ckpt] step{step} ({size_mb:.0f}MB) loss={loss:.4f}", flush=True)
    upload_hf(path, f"checkpoints/fractus_1b_step{step}.pt")
    upload_hf(path, "checkpoints/fractus_1b_latest.pt")

    # Garbage-collect old step checkpoints (keep the last `keep_last`).
    import glob
    step_ckpts = sorted(glob.glob(os.path.join(ckpt_dir, "fractus_1b_step*.pt")),
                        key=lambda p: int(p.split("step")[-1].split(".")[0]))
    for old in step_ckpts[:-keep_last]:
        os.remove(old)
        print(f"  [disk] Removed old step ckpt {os.path.basename(old)}", flush=True)


def chunked_cross_entropy(model, hidden, target, vocab, chunk_positions):
    """Compute lm_head + cross-entropy by chunks of positions to avoid

    materializing the full (B, L, vocab) tensor.



    hidden:  (B, L, d_model) β€” the final hidden states (already through all blocks).

    target:  (B, L) β€” next-token ids.

    Returns scalar loss (averaged over all positions).



    Processes positions in chunks of `chunk_positions`. Each chunk materializes

    only (B, chunk, vocab) β€” keeping VRAM low so batch can grow 4-8x.

    """
    B, L, _ = hidden.shape
    total_loss = 0.0
    n = 0
    # Detach hidden from the chunk loop's graph accumulation β€” we sum losses
    # and backward once at the end. Each chunk's logits share the same hidden,
    # so the gradient flows correctly through hidden to all blocks.
    losses = []
    for s in range(0, L, chunk_positions):
        e = min(s + chunk_positions, L)
        h_chunk = hidden[:, s:e]                      # (B, C, D)
        logits_chunk = model.lm_head(h_chunk)         # (B, C, vocab)
        tgt_chunk = target[:, s:e]                     # (B, C)
        l = F.cross_entropy(
            logits_chunk.reshape(-1, vocab),
            tgt_chunk.reshape(-1),
            reduction="sum",
        )
        losses.append(l)
        n += (e - s) * B
    total = torch.stack(losses).sum() / n
    return total


def main():
    parser = argparse.ArgumentParser(description="Fractus-1B Cloud Training")
    parser.add_argument("--epochs", type=int, default=20)
    parser.add_argument("--seq-len", type=int, default=64,
                       help="Longer seq = better context (GPU can handle it)")
    parser.add_argument("--batch-size", type=int, default=8,
                       help="Batch size (GPU parallelism)")
    parser.add_argument("--lr", type=float, default=1e-4)
    parser.add_argument("--corpus", type=str, default=None,
                       help="Path to corpus. If not found, builds it.")
    parser.add_argument("--resume", type=str, default=None,
                       help="Checkpoint to resume from")
    parser.add_argument("--upload-every", type=int, default=1,
                       help="Upload checkpoint every N epochs")
    parser.add_argument("--log-every", type=int, default=500,
                       help="Log every N steps")
    parser.add_argument("--compile", dest="compile", action="store_true",
                       default=True,
                       help="Enable torch.compile (default ON on GPU)")
    parser.add_argument("--no-compile", dest="compile", action="store_false",
                       help="Disable torch.compile, use eager")
    parser.add_argument("--chunk-ce", type=int, default=0,
                       help="Chunk positions for CE (0=disabled, 8=recommended). "
                            "Avoids materializing full (B,L,vocab) tensor β†’ bigger batch.")
    parser.add_argument("--triton-ce", dest="triton_ce", action="store_true",
                       default=True,
                       help="Use Triton fused linear+CE kernel (default ON on GPU). "
                            "Auto self-test; falls back if unavailable.")
    parser.add_argument("--no-triton-ce", dest="triton_ce", action="store_false")
    parser.add_argument("--save-every", type=int, default=10000,
                       help="Save+upload checkpoint every N steps (crash recovery). "
                            "Default 10000 = ~6%% of a 1.76B-token epoch at batch 512.")
    args = parser.parse_args()

    # Detect device.
    if torch.cuda.is_available():
        device = torch.device("cuda")
        gpu_name = torch.cuda.get_device_name(0)
        vram = torch.cuda.get_device_properties(0).total_memory / 1e9
        print(f"GPU: {gpu_name} ({vram:.1f} GB VRAM)", flush=True)
        torch.backends.cudnn.benchmark = True
    else:
        device = torch.device("cpu")
        print("WARNING: No GPU detected. Running on CPU.", flush=True)

    torch.manual_seed(42)
    num_threads = os.cpu_count() or 4
    torch.set_num_threads(num_threads)
    print(f"Threads: {num_threads}", flush=True)

    # Load or build corpus.
    script_dir = os.path.dirname(os.path.abspath(__file__))
    project_dir = os.path.dirname(script_dir)
    
    if args.corpus:
        corpus_path = args.corpus
    else:
        # Try communication corpus first, then ultimate, then mega.
        for name in ["communication_corpus.pt", "ultimate_corpus.pt", "mega_corpus.pt"]:
            p = os.path.join(project_dir, "data", name)
            if os.path.exists(p):
                corpus_path = p
                break
        else:
            # Build it.
            print("No corpus found. Building communication corpus...", flush=True)
            import subprocess
            subprocess.run([sys.executable, os.path.join(script_dir, "build_communication_corpus.py")], check=True)
            corpus_path = os.path.join(project_dir, "data", "communication_corpus.pt")

    print(f"Loading corpus: {corpus_path}", flush=True)
    tokens = torch.load(corpus_path, weights_only=False).long()
    print(f"Corpus: {len(tokens):,} tokens", flush=True)

    # Build model.
    print("Building Fractus-1B...", flush=True)
    model = Fractus1B(
        vocab_size=50257, d_model=768, n_layers=8, n_heads=12, d_head=64,
        n_levels=2, n_experts=64, top_k=2, expert_d_ff=1024, siren_rank=16,
        max_seq_len=args.seq_len,
    ).to(device)
    n = model.n_params()
    cap = model.n_effective_capacity()
    print(f"  Params: {n:,} ({n/1e6:.0f}M)", flush=True)
    print(f"  Capacity: {cap:,} ({cap/1e9:.2f}B)", flush=True)
    print(f"  RAM: {n*4/1e9:.1f}GB", flush=True)

    # torch.compile β€” now that the MoE is vectorized (no dynamic control flow),
    # compile can fuse kernels. Cache limit raised to handle the 64 expert guards.
    # GUARD: wrapped in try/except, falls back to eager on any failure.
    use_compiled = False
    if args.compile and device.type == "cuda":
        try:
            import torch._dynamo as dyn
            dyn.config.cache_size_limit = 256
            dyn.config.accumulated_cache_size_limit = 512
            model = torch.compile(model, mode="reduce-overhead", dynamic=False)
            use_compiled = True
            print("  torch.compile: ON (mode=reduce-overhead)", flush=True)
        except Exception as e:
            print(f"  torch.compile: FAILED ({type(e).__name__}), using eager", flush=True)
    else:
        print(f"  torch.compile: OFF ({'disabled by --no-compile' if not args.compile else 'CPU device'})", flush=True)

    # Triton fused kernel β€” runs self-test before use. Falls back to eager/chunk-ce.
    use_triton_ce = False
    if device.type == "cuda" and _HAS_TRITON_IMPORT and args.triton_ce:
        try:
            ok = triton_self_test()
            use_triton_ce = bool(ok)
            print(f"  triton fused-CE: {'ON' if use_triton_ce else 'OFF (self-test failed)'}", flush=True)
        except Exception as e:
            print(f"  triton fused-CE: FAILED ({type(e).__name__})", flush=True)
    else:
        print(f"  triton fused-CE: OFF (cuda={device.type=='cuda'}, import={_HAS_TRITON_IMPORT})", flush=True)

    # Resume if specified.
    start_epoch = 0
    start_step = 0
    if args.resume:
        print(f"Resuming from: {args.resume}", flush=True)
        ckpt = torch.load(args.resume, weights_only=False, map_location=device)
        model.load_state_dict(ckpt["model_state"])
        start_epoch = ckpt.get("epoch", 0)
        # Restore the global step counter so checkpoint names + logs use the
        # true global step, not a local counter that restarts at 0 on resume.
        # This was a bug: resuming from step 140000 reset the counter to 0,
        # causing new checkpoints to overwrite old ones on HF by name.
        start_step = ckpt.get("step", start_epoch * n_steps)
        # Restore optimizer state if present (so Adam moments are preserved).
        if "optimizer_state" in ckpt:
            try:
                opt.load_state_dict(ckpt["optimizer_state"])
                print(f"  Optimizer state restored", flush=True)
            except Exception as e:
                print(f"  Optimizer state load failed: {e}", flush=True)
        print(f"  Resumed from epoch {start_epoch}, step {start_step}, loss={ckpt.get('loss','?')}", flush=True)

    # Optimizer + scheduler.
    opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01)
    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=args.epochs, eta_min=1e-6)
    if start_epoch > 0:
        for _ in range(start_epoch):
            sched.step()

    # AMP for GPU.
    use_amp = device.type == "cuda"
    scaler = torch.cuda.amp.GradScaler() if use_amp else None

    tok = FractusTokenizer.gpt2_compatible()
    seq = args.seq_len
    batch_size = args.batch_size
    ckpt_dir = os.path.join(project_dir, "checkpoints")
    n_steps = len(tokens) // seq // batch_size

    print(f"\nTraining {args.epochs} epochs", flush=True)
    print(f"  seq_len={seq}, batch_size={batch_size}, lr={args.lr}", flush=True)
    print(f"  {n_steps:,} steps/epoch", flush=True)
    print(f"  AMP: {'ON' if use_amp else 'OFF'}", flush=True)
    print(f"  HF upload: every {args.upload_every} epochs β†’ {HF_REPO}", flush=True)
    print("=" * 70, flush=True)

    initial_loss = None
    step = start_step  # global step counter (preserved across resumes)

    for epoch in range(start_epoch, args.epochs):
        model.train()
        t0 = time.perf_counter()
        ep_loss = 0.0
        ep_n = 0

        # Create batches.
        for batch_start in range(0, len(tokens) - seq * batch_size - 1, seq * batch_size):
            # Build batch.
            inp_list = []
            tgt_list = []
            for b in range(batch_size):
                offset = batch_start + b * seq
                inp_list.append(tokens[offset:offset + seq])
                tgt_list.append(tokens[offset + 1:offset + seq + 1])
            inp = torch.stack(inp_list).to(device)
            tgt = torch.stack(tgt_list).to(device)

            opt.zero_grad()
            if use_amp:
                with torch.cuda.amp.autocast(dtype=torch.bfloat16):
                    if use_triton_ce:
                        # Triton fused path: skip lm_head in model, kernel does it.
                        model._return_hidden = True
                        hidden, aux = model(inp)
                        model._return_hidden = False
                        ce = fused_linear_cross_entropy(hidden, model.lm_head.weight, tgt)
                    elif args.chunk_ce > 0:
                        model._return_hidden = True
                        hidden, aux = model(inp)
                        ce = chunked_cross_entropy(model, hidden, tgt, 50257, args.chunk_ce)
                        model._return_hidden = False
                    else:
                        logits, aux = model(inp)
                        ce = F.cross_entropy(logits.reshape(-1, 50257), tgt.reshape(-1))
                    # Clip aux (load-balance loss) to prevent the divergence seen at step 149000.
                    # When all tokens route to one expert, lb_loss can spike to 5-15 and kill the
                    # main CE gradient. Cap its contribution at 0.001.
                    aux_clamped = torch.clamp(aux, max=1.0)
                    loss = ce + 0.001 * aux_clamped
                loss.backward()
            else:
                if use_triton_ce:
                    model._return_hidden = True
                    hidden, aux = model(inp)
                    model._return_hidden = False
                    ce = fused_linear_cross_entropy(hidden, m.lm_head.weight, tgt)
                elif args.chunk_ce > 0:
                    model._return_hidden = True
                    hidden, aux = model(inp)
                    ce = chunked_cross_entropy(model, hidden, tgt, 50257, args.chunk_ce)
                    model._return_hidden = False
                else:
                    logits, aux = model(inp)
                    ce = F.cross_entropy(logits.reshape(-1, 50257), tgt.reshape(-1))
                aux_clamped = torch.clamp(aux, max=1.0)
                loss = ce + 0.001 * aux_clamped
                loss.backward()
            # Skip the step if loss became NaN/inf (defensive β€” should not happen post-clamp).
            if not torch.isfinite(loss):
                opt.zero_grad()
                continue
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            opt.step()

            loss_val = loss.item()
            ep_loss += loss_val
            ep_n += 1
            step += 1
            if initial_loss is None:
                initial_loss = loss_val
            if step % args.log_every == 0:
                elapsed = time.perf_counter() - t0
                sps = (step - (epoch * n_steps)) / max(elapsed, 1e-6)
                ppl = math.exp(min(loss_val, 20))
                print(f"  E{epoch} S{step:>7}/{n_steps*(epoch+1):>7} "
                      f"loss={loss_val:.4f} ppl={ppl:.1f} aux={aux.item():.4f} "
                      f"{sps:.1f}step/s", flush=True)
            # Mid-epoch checkpoint for crash recovery (every save_every steps).
            if args.save_every > 0 and step % args.save_every == 0:
                save_step_checkpoint(
                    model, opt, step, epoch, loss_val,
                    {"seq_len": seq, "batch_size": batch_size,
                     "lr": args.lr, "corpus": args.corpus},
                    ckpt_dir, keep_last=2,
                )

        # End of epoch.
        avg = ep_loss / max(ep_n, 1)
        elapsed = time.perf_counter() - t0
        print(f"\n[EPOCH {epoch}] avg_loss={avg:.4f} ppl={math.exp(min(avg,20)):.1f} "
              f"time={elapsed/60:.1f}min", flush=True)
        save_and_upload(model, opt, epoch, avg, initial_loss,
                        {"seq_len": seq, "batch_size": batch_size,
                         "lr": args.lr, "corpus": args.corpus}, ckpt_dir)
        sched.step()

        # Sample generation to monitor quality.
        try:
            model.eval()
            with torch.no_grad():
                prompt = "def fibonacci"
                ids = tok.encode(prompt)
                x = torch.tensor([ids], device=device)
                out = []
                for _ in range(60):
                    lg, _ = model(x)
                    nxt = lg[0, -1].argmax().unsqueeze(0).unsqueeze(0)
                    x = torch.cat([x, nxt], dim=1)
                    out.append(nxt.item())
                    if nxt.item() == 50256:
                        break
                print(f"[SAMPLE] {(prompt + tok.decode(out))[:300]}\n", flush=True)
            model.train()
        except Exception as e:
            print(f"  [sample] failed: {e}", flush=True)

    print("\n" + "=" * 70, flush=True)
    print(f"TRAINING COMPLETE. Final loss={avg:.4f}", flush=True)
    print(f"Loss reduction: {initial_loss:.4f} -> {avg:.4f} "
          f"({(initial_loss-avg)/initial_loss*100:.1f}% improvement)", flush=True)


if __name__ == "__main__":
    main()