#!/usr/bin/env python """Train a Fractus CTE from zero on Atomizer ids. Source architecture: HF thefinalboss/fractus-cte (ContinuousThoughtEngine). This script never loads an x8 / GPT-2 checkpoint. vocab 50257 weights cannot map onto vocab 266. START_TOKEN is 0 because the run is new — that exception does not apply to the existing x8 resume. Usage (smoke, CPU): python scripts/train_atom_from_scratch.py --scale smoke --steps 30 Usage (1B config, fresh, on a pod — do not point this at x8run): python scripts/train_atom_from_scratch.py --scale 1b --corpus data/atom_corpus.i16 """ from __future__ import annotations import argparse import json import os import sys import time sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import numpy as np import torch import torch.nn.functional as F from fractus.atom_tokenizer import VOCAB_SIZE, AtomFractusTokenizer from fractus.continuous_engine import ContinuousThoughtEngine from fractus.grow import grow_atom from fractus.kuramoto_fix import apply_kuramoto_routing_fix from fractus.train.ar_loss import ss_prob_at from fractus.train.v4_step import should_ss, v4_forward_losses, v4_ss_pass SCALE = { "smoke": dict( d_model=64, n_heads=4, d_head=16, n_levels=1, n_oscillators=4, coupling_rank=2, n_experts=4, top_k=2, expert_d_ff=64, siren_rank=8, n_layers=2, ), "1b": dict( d_model=1280, n_heads=20, d_head=64, n_levels=2, n_oscillators=16, coupling_rank=8, n_experts=128, top_k=2, expert_d_ff=2048, siren_rank=64, n_layers=16, ), } SMOKE_TEXT = ( "Fractus pense en continu. L'Atomizer coupe le flux en spans d'octets, " "pas en BPE. Bonjour Philippe. 2+2=4. def tick(): return h\n" ) * 8 def load_stream(path: str | None, tok: AtomFractusTokenizer): if not path: ids, feat = tok.encode_with_features(SMOKE_TEXT) return torch.tensor(ids, dtype=torch.long), feat if path.endswith(".i16"): arr = np.fromfile(path, dtype=np.int16) elif path.endswith(".npy"): arr = np.load(path, mmap_mode="r") else: raise SystemExit(f"unsupported corpus: {path}") if arr.size and int(arr.max()) >= VOCAB_SIZE: raise SystemExit("corpus id >= vocab 266 — this is not an Atom stream") return torch.as_tensor(np.asarray(arr, dtype=np.int64)), None def main(): ap = argparse.ArgumentParser() ap.add_argument("--scale", choices=sorted(SCALE), default="smoke") ap.add_argument("--corpus", default=None) ap.add_argument("--steps", type=int, default=40) ap.add_argument("--seq-len", type=int, default=32) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--out", default="checkpoints/atom_scratch") ap.add_argument("--max-span-bytes", type=int, default=32) ap.add_argument("--pack-mode", default="linguistic") ap.add_argument("--grow-at", type=int, default=-1, help="step at which the body grows by one layer") ap.add_argument("--ss-rate", type=float, default=0.0) args = ap.parse_args() tok = AtomFractusTokenizer( max_span_bytes=args.max_span_bytes, pack_mode=args.pack_mode ) ids, feat = load_stream(args.corpus, tok) if ids.numel() < args.seq_len + 1: raise SystemExit("corpus shorter than seq-len+1") cfg = dict(SCALE[args.scale]) engine = ContinuousThoughtEngine(vocab_size=VOCAB_SIZE, **cfg) routing = apply_kuramoto_routing_fix(engine, log=lambda *_: None) n_params = sum(p.numel() for p in engine.parameters()) opt = torch.optim.AdamW(engine.parameters(), lr=args.lr, weight_decay=0.01) engine.train() losses = [] t0 = time.time() n = ids.numel() grew_at = None for step in range(args.steps): if step == args.grow_at: engine = grow_atom(engine, {"n_layers": engine.n_layers + 1}) opt = torch.optim.AdamW(engine.parameters(), lr=args.lr, weight_decay=0.01) grew_at = step start = (step * args.seq_len) % (n - args.seq_len - 1) chunk = ids[start:start + args.seq_len].unsqueeze(0) target = ids[start + 1:start + 1 + args.seq_len] feat_chunk = None if feat is not None: feat_chunk = feat[start:start + args.seq_len].unsqueeze(0) engine.reset_thought(batch_size=1) loss, extras = v4_forward_losses(engine, chunk, target, feat=feat_chunk) if args.ss_rate and should_ss(args.ss_rate): ss_loss, _ = v4_ss_pass( engine, chunk, target, extras["h"], ss_prob=ss_prob_at(step * args.seq_len), ) loss = loss + ss_loss opt.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(engine.parameters(), 1.0) opt.step() losses.append(float(extras["ce"].detach())) if step % max(1, args.steps // 5) == 0 or step == args.steps - 1: print( f"step {step} ce {losses[-1]:.4f} lb {float(extras['lb'].detach()):.4f} " f"repeat {float(extras['repeat'].detach()):.4f}", flush=True, ) os.makedirs(args.out, exist_ok=True) ckpt = os.path.join(args.out, f"fractus_atom_{args.scale}.pt") torch.save( { "model_state": engine.state_dict(), "config": {**cfg, "vocab_size": VOCAB_SIZE, "scale": args.scale}, "tokenizer": { "version": tok.VERSION, "vocab_size": VOCAB_SIZE, "max_span_bytes": args.max_span_bytes, "pack_mode": args.pack_mode, }, "steps": args.steps, "start_token": 0, "parent": "hf:thefinalboss/fractus-cte", "loop": "v4_forward_losses", "kuramoto_fix": routing, "grew_at": grew_at, "note": "from scratch. not an x8 resume.", }, ckpt, ) summary = { "ckpt": ckpt, "params": n_params, "vocab_size": VOCAB_SIZE, "steps": args.steps, "loss_first": losses[0], "loss_last": losses[-1], "seconds": round(time.time() - t0, 2), "tokens_seen": args.steps * args.seq_len, } with open(os.path.join(args.out, "scratch_summary.json"), "w", encoding="utf-8") as handle: json.dump(summary, handle, indent=2) print(json.dumps(summary)) if __name__ == "__main__": main()