""" Train Quazimoto-LM with the SpikeWhale tokenizer and the family's AdamW regime. Tokenizer: the custom SpikeWhale length-max byte tokenizer (tokenizer.json + spike_tokenizer.py) bundled in this package -- self-contained, no external folder. Data is tokenized fresh in Python from a local UTF-8 .txt (or a synthetic fallback). Optimizer matches the family's logged-stable choice (LEVEL6_PLAN): AdamW, peak lr 3e-4, linear warmup, cosine decay to min_lr_frac*peak (0.3 floor), grad-clip 1.0. MuonEq-R is deliberately NOT used here -- it caused weight runaway in this family. """ import argparse, time, math, os, sys import numpy as np import torch from model import QuazimotoLM, QuazimotoConfig # SpikeWhale tokenizer is bundled in this package (self-contained, no level6 needed). PKG_DIR = os.path.dirname(os.path.abspath(__file__)) def load_tokenizer(tok_dir): sys.path.insert(0, tok_dir) from spike_tokenizer import SpikeTokenizer tok = SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json")) return tok def load_ids(path, tok): if path and os.path.exists(path): with open(path, "r", encoding="utf-8", errors="replace") as f: text = f.read() else: text = ("the quazimoto oscillator learns to synchronize phases across rings. " "coupled clocks fall into step when coupling is strong enough. ") * 4000 ids = np.array(tok.encode(text, add_special_tokens=False), dtype=np.int64) n = int(0.9 * len(ids)) return ids[:n], ids[n:] def get_batch(ids, bs, block, device): ix = np.random.randint(0, len(ids) - block - 1, size=bs) x = np.stack([ids[i:i + block] for i in ix]) y = np.stack([ids[i + 1:i + 1 + block] for i in ix]) return (torch.from_numpy(x).to(device), torch.from_numpy(y).to(device)) def _doc_text(ex): """Pull the text field from a dataset example, tolerating naming differences across datasets (text / content / raw_content / document), else the first non-trivial string value.""" for k in ("text", "content", "raw_content", "document"): v = ex.get(k) if isinstance(v, str) and v.strip(): return v for v in ex.values(): if isinstance(v, str) and len(v.strip()) > 1: return v return None def _stream_docs(path, config, split): """Infinite generator of non-empty doc strings from a streaming HF dataset. Re-opens the stream when a shard pass is exhausted so training never runs dry. (Mirrors the level6 train_snn_stream._stream_docs convention.)""" from datasets import load_dataset while True: ds = load_dataset(path, name=config, split=split, streaming=True) any_yielded = False for ex in ds: text = _doc_text(ex) if text: any_yielded = True yield text if not any_yielded: raise RuntimeError(f"stream {path}:{config} yielded no usable text") def _blend_sources(args): """Return the active (label, path, config, weight) sources for the blend. Default mix: 35% Ultra-FineWeb-L3 / 25% FineWeb-Edu / 25% FineMath / 15% PretrainNew.""" srcs = [ ("ultra", args.ultra_dataset, args.ultra_config, args.ultra_frac), ("edu", args.fineweb_dataset, args.fineweb_config, args.edu_frac), ("math", args.math_dataset, args.math_config, args.math_frac), ("pretrain", args.pretrain_dataset, args.pretrain_config, args.pretrain_frac), ] return [s for s in srcs if s[3] > 0] def doc_generator(args): """Blend multiple HF datasets by per-document probability (weights normalized).""" import random rng = random.Random(args.seed) srcs = _blend_sources(args) gens = [_stream_docs(p, c, args.split) for (_, p, c, _) in srcs] weights = [s[3] for s in srcs] tot = sum(weights) cum, acc = [], 0.0 for w in weights: # cumulative thresholds for sampling acc += w / tot cum.append(acc) while True: r = rng.random() gi = next(k for k, c in enumerate(cum) if r <= c) yield next(gens[gi]) + "\n" def stream_batches(tok, args, device): """Infinite [B,T] next-token batches from the streamed blend via a token buffer. Docs are EOS-separated and packed; each batch is B contiguous T+1 windows.""" eos = tok.eos_token_id sep = [eos] if eos is not None else [] need = args.batch * (args.block + 1) buf, docs = [], doc_generator(args) srcs = _blend_sources(args) tot = sum(s[3] for s in srcs) blend = " / ".join(f"{int(round(100*s[3]/tot))}% {s[1]}" for s in srcs) print(f"streaming blend: {blend}") while True: while len(buf) < need: buf.extend(tok.encode(next(docs), add_special_tokens=False)) buf.extend(sep) chunk = np.array(buf[:need], dtype=np.int64).reshape(args.batch, args.block + 1) del buf[:need] x = torch.from_numpy(chunk[:, :-1]).to(device) y = torch.from_numpy(chunk[:, 1:]).to(device) yield x, y def save_ckpt(model, tok_vocab, step, path, opt=None): if not path: return os.makedirs(os.path.dirname(path) or ".", exist_ok=True) ckpt = {"model": model.state_dict(), "family_config": model.family_config, "vocab_size": tok_vocab, "step": step} if opt is not None: ckpt["optim"] = opt.state_dict() # so --resume continues the optimizer too torch.save(ckpt, path) print(f" saved checkpoint -> {path} (step {step})") def find_latest_ckpt(out_path): """Auto-locate a checkpoint to resume from: prefer the exact --out file, else the most recently modified *.pt in the same folder. Returns None if none.""" if out_path and os.path.isfile(out_path): return out_path folder = os.path.dirname(out_path) or "." if not os.path.isdir(folder): return None pts = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(".pt")] return max(pts, key=os.path.getmtime) if pts else None def lr_at(step, peak, warmup, total, min_frac): if step < warmup: return peak * step / max(warmup, 1) if step >= total: return peak * min_frac prog = (step - warmup) / max(total - warmup, 1) return peak * (min_frac + (1 - min_frac) * 0.5 * (1 + math.cos(math.pi * prog))) def main(): p = argparse.ArgumentParser() p.add_argument("--data", default="") p.add_argument("--tok-dir", default=PKG_DIR, help="dir with bundled tokenizer.json + spike_tokenizer.py") p.add_argument("--steps", type=int, default=200) p.add_argument("--batch", type=int, default=8) p.add_argument("--block", type=int, default=256) p.add_argument("--lr", type=float, default=3e-4) # family peak p.add_argument("--warmup", type=int, default=1000) p.add_argument("--min-lr-frac", type=float, default=0.3) # family floor (~9e-5) p.add_argument("--weight-decay", type=float, default=0.01) p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") p.add_argument("--use-hrm", action="store_true") p.add_argument("--use-moe", action="store_true") p.add_argument("--use-mtp", action="store_true") p.add_argument("--mtp-layers", type=int, default=4, help="MTP draft-head depth (spec-decode)") p.add_argument("--use-jepa", action="store_true") p.add_argument("--use-rings", action="store_true") p.add_argument("--use-ring-controllers", action="store_true") p.add_argument("--use-ring-specialists", action="store_true", help="per-ring MoE memory specialists with test-time input/output stores") p.add_argument("--use-fractal-phase-seed", action="store_true", help="seed oscillator phases from each token's Mandelbrot orbit (gated)") p.add_argument("--out", default=os.path.join(PKG_DIR, "chkpt", "quazimoto.pt")) # checkpoint file path p.add_argument("--ckpt-every", type=int, default=250) p.add_argument("--resume", action="store_true", help="auto-find the latest checkpoint in the --out folder and continue training") # streaming blend (HF datasets); off by default so local runs stay dependency-free. # Default 4-way mix: 35% Ultra-FineWeb-L3 / 25% FineWeb-Edu / 25% FineMath / 15% PretrainNew. p.add_argument("--stream", action="store_true", help="stream the weighted dataset blend") p.add_argument("--ultra-frac", type=float, default=0.35) p.add_argument("--ultra-dataset", default="openbmb/Ultra-FineWeb-L3") p.add_argument("--ultra-config", default="Ultra-FineWeb-L3-en-Multi-Style-Synthetic") p.add_argument("--edu-frac", type=float, default=0.25) p.add_argument("--fineweb-dataset", default="HuggingFaceFW/fineweb-edu") p.add_argument("--fineweb-config", default="sample-10BT") p.add_argument("--math-frac", type=float, default=0.25) p.add_argument("--math-dataset", default="HuggingFaceTB/finemath") p.add_argument("--math-config", default="finemath-4plus") p.add_argument("--pretrain-frac", type=float, default=0.15) p.add_argument("--pretrain-dataset", default="nvidia/Nemotron-Pretraining-Specialized-v1.1") p.add_argument("--pretrain-config", default="Nemotron-Pretraining-Formal-Logic") p.add_argument("--split", default="train") p.add_argument("--seed", type=int, default=0) p.add_argument("--amp", action="store_true", help="bf16 autocast (biggest single-GPU speedup; also frees memory)") args = p.parse_args() # single-GPU speed: TF32 matmuls (free, ~fp32 accuracy on Ampere+/Blackwell). if args.device == "cuda": torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True tok = load_tokenizer(args.tok_dir) cfg = QuazimotoConfig(vocab_size=tok.vocab_size, block_size=args.block, use_hrm=args.use_hrm, use_moe=args.use_moe, use_mtp=args.use_mtp, mtp_layers=args.mtp_layers, use_jepa=args.use_jepa, use_rings=args.use_rings, use_ring_controllers=args.use_ring_controllers, use_ring_specialists=args.use_ring_specialists, use_fractal_phase_seed=args.use_fractal_phase_seed) model = QuazimotoLM(cfg).to(args.device) print(f"tokenizer: SpikeWhale (vocab {tok.vocab_size})") if args.stream: batches = stream_batches(tok, args, args.device) # infinite blended stream train = None else: train, _ = load_ids(args.data, tok) print(f"data: {len(train)} local tokens | device {args.device}") opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=args.weight_decay, fused=(args.device == "cuda")) # fused CUDA optimizer step start_step = 1 if args.resume: ckpt_path = find_latest_ckpt(args.out) if ckpt_path is None: print(f"--resume: no checkpoint found near {args.out}; starting fresh.") else: ck = torch.load(ckpt_path, map_location=args.device, weights_only=False) miss, unexp = model.load_state_dict(ck["model"], strict=False) if miss: print(f" [resume warn] missing keys: {len(miss)} (e.g. {miss[:2]})") if unexp: print(f" [resume warn] unexpected keys: {len(unexp)} (e.g. {unexp[:2]})") if "optim" in ck: try: opt.load_state_dict(ck["optim"]) except ValueError as e: print(f" [resume warn] optimizer state not restored ({e}); using fresh optimizer.") start_step = int(ck.get("step", 0)) + 1 print(f"resumed from {ckpt_path} at step {ck.get('step')} -> continuing at {start_step}") if start_step > args.steps: print(f" already at/past --steps ({args.steps}); nothing to do.") model.train() t0 = time.time() for step in range(start_step, args.steps + 1): lr = lr_at(step, args.lr, args.warmup, args.steps, args.min_lr_frac) for g in opt.param_groups: g["lr"] = lr x, y = (next(batches) if args.stream else get_batch(train, args.batch, args.block, args.device)) # bf16 autocast: ~1.5-2x on Blackwell tensor cores. bf16 keeps fp32 range so # no grad scaler is needed, and it halves activation memory (room for bigger # --batch). Params stay fp32; the oscillator trig/erf math is safe in bf16. with torch.autocast("cuda", dtype=torch.bfloat16, enabled=args.amp and args.device == "cuda"): _, loss, aux = model(x, y) total = loss + sum(aux.values()) # add active trait aux losses opt.zero_grad(set_to_none=True) total.backward() gn = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) if not torch.isfinite(gn): # family grad-finiteness guard print(f"step {step}: non-finite grad norm, skipping") opt.zero_grad(set_to_none=True) continue opt.step() if step % 10 == 0 or step == 1: dt = time.time() - t0 extra = " ".join(f"{k} {float(v):.3f}" for k, v in aux.items()) bank = getattr(model, "ring_bank", None) surp = f" surprise {float(bank.last_surprise):.3f}" if bank is not None else "" print(f"step {step:4d} | loss {loss.item():.3f} | " f"bpt {loss.item()/math.log(2):.3f} | lr {lr:.2e} | {extra}{surp} | {dt:.1f}s") if args.ckpt_every and step % args.ckpt_every == 0: save_ckpt(model, tok.vocab_size, step, args.out, opt) save_ckpt(model, tok.vocab_size, args.steps, args.out, opt) print("done.") if __name__ == "__main__": main()