fractus-cte-atom / scripts /train_atom_from_scratch.py
thefinalboss's picture
v4 loop, Siren synced on grow, unique40 probe, Atom features.
2841efa verified
Raw History Blame Contribute Delete
6.47 kB
#!/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()