fractus-cte-atom / scripts /fast4gpu_atom.py
thefinalboss's picture
Compile the pure core. B=16 hits 5896 ids/s.
9a58001 verified
Raw History Blame Contribute Delete
19.7 kB
#!/usr/bin/env python3
"""Fractus CTE-Atom trainer — the 1B loop of fast4gpu_boost_v4, made Atom-safe.
Why a new file: boost_v4 cannot train this fork.
- It loads an x8 checkpoint by default and slices the 50257 head to 266 rows
(BPE rows read as bytes), with strict=False.
- It refuses anything that is not a .npy int32 shard, so atom_corpus.i16 is rejected.
- It sets the gate temperature but not omega x4 (partial Kuramoto fix).
- It calls eng.set_p0_routing() and eng.routing_stats(), which do not exist on this
engine, and reads probe keys (gate_go, mean_unique, echo_frac, head) that
unique40_probe no longer returns. It stops before the first step.
- Its tok/s counts B*SEQ once per step while SS_RATE=1.0 runs two forward+backward.
What this file does instead:
- No CKPT_IN -> fresh engine, vocab 266, full apply_kuramoto_routing_fix, start at id 0.
- CKPT_IN set -> Atom checkpoints only: head and embed must have 266 rows, strict load,
no slicing, Kuramoto fix NOT re-applied (omega x4 would compound), START_TOKEN required
and must not be 0.
- Corpus: .i16 (raw int16, as written by build_atom_corpus.py) or .npy. Every id is
checked to be in 0..265 before step 1.
- Multi-GPU: N_GPU independent processes, each on a contiguous slice of the corpus.
- Throughput is reported in Atom ids, split: TF ids/s, SS ids/s, probe time, wall.
It never opens a pod and never downloads anything.
CPU smoke (small body):
SCALE=smoke CORPUS=data/atom_corpus.i16 BATCH=1 SEQ=32 MAX_STEPS=20 \\
python -u scripts/fast4gpu_atom.py
One GPU benchmark (1B shape, fresh):
CUDA_VISIBLE_DEVICES=0 GPU_ID=0 N_GPU=1 CORPUS=data/atom_corpus.i16 \\
BATCH=4 SEQ=128 FRACTUS_ATTN_IMPL=chunked BLOCK_CKPT=1 MAX_STEPS=200 \\
python -u scripts/fast4gpu_atom.py
"""
from __future__ import annotations
import json
import os
import sys
import time
from pathlib import Path
import numpy as np
import torch
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
os.environ.setdefault("FRACTUS_ATTN_IMPL", "chunked")
from fractus.atom_tokenizer import VOCAB_SIZE # noqa: E402
from fractus.continuous_engine import ContinuousThoughtEngine # noqa: E402
from fractus.generate_aligned import unique40_probe # noqa: E402
from fractus.kuramoto_fix import apply_kuramoto_routing_fix # noqa: E402
from fractus.train.ar_loss import ss_prob_at # noqa: E402
from fractus.train.v4_step import ( # noqa: E402
restore_carry, should_ss, snapshot_carry, v4_forward_losses, v4_ss_pass,
)
assert VOCAB_SIZE == 266, VOCAB_SIZE
SCALES = {
"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,
),
}
CFG_KEYS = tuple(SCALES["1b"].keys())
# ---------------------------------------------------------------- guards
def open_corpus(path: str):
"""Memmap an Atom stream. Refuses anything that is not ids in 0..265."""
p = str(path)
if p.endswith(".i16"):
arr = np.memmap(p, dtype=np.int16, mode="r")
elif p.endswith(".npy"):
arr = np.load(p, mmap_mode="r")
if arr.ndim != 1 or arr.dtype.kind not in "iu":
raise SystemExit(f"corpus {p}: need a 1-D integer array, got {arr.dtype} {arr.shape}")
else:
raise SystemExit(f"corpus {p}: need .i16 or .npy")
n = int(arr.shape[0])
if n == 0:
raise SystemExit(f"corpus {p}: empty")
step = 1 << 24
lo, hi = 1 << 30, -1
for i in range(0, n, step):
part = np.asarray(arr[i:i + step])
lo, hi = min(lo, int(part.min())), max(hi, int(part.max()))
if lo < 0 or hi >= VOCAB_SIZE:
raise SystemExit(
f"corpus {p}: ids span {lo}..{hi}, Atom ids must be 0..{VOCAB_SIZE - 1}. "
"This is not an Atom stream (GPT-2 shards are refused)."
)
return arr, n
def check_atom_state(sd: dict, where: str) -> None:
"""Refuse any state dict whose embed or head is not exactly 266 rows."""
for key in ("observe.weight", "output_head.weight"):
if key not in sd:
raise SystemExit(f"{where}: missing {key}, not a ContinuousThoughtEngine checkpoint")
rows = int(sd[key].shape[0])
if rows != VOCAB_SIZE:
raise SystemExit(
f"{where}: {key} has {rows} rows, need {VOCAB_SIZE}. "
"x8 / GPT-2 weights do not map onto Atom ids. Not loading, not slicing."
)
def build_engine(scale: str, ckpt_in: str | None, log):
"""Return (engine, cfg, birth) where birth records how the weights came to be."""
if not ckpt_in:
cfg = dict(SCALES[scale])
eng = ContinuousThoughtEngine(vocab_size=VOCAB_SIZE, **cfg)
fix = apply_kuramoto_routing_fix(eng, log=log)
return eng, cfg, {"fresh": True, "kuramoto_fix": fix, "parent": "hf:thefinalboss/fractus-cte"}
ck = torch.load(ckpt_in, map_location="cpu", weights_only=False)
sd = ck.get("model_state", ck)
sd = {(k[10:] if k.startswith("_orig_mod.") else k): v for k, v in sd.items()}
check_atom_state(sd, ckpt_in)
ck_cfg = ck.get("config", {}) if isinstance(ck, dict) else {}
if int(ck_cfg.get("vocab_size", VOCAB_SIZE)) != VOCAB_SIZE:
raise SystemExit(f"{ckpt_in}: config vocab_size={ck_cfg.get('vocab_size')}, need {VOCAB_SIZE}")
missing = [k for k in CFG_KEYS if k not in ck_cfg]
if missing:
raise SystemExit(f"{ckpt_in}: config lacks {missing}; cannot rebuild the body exactly")
cfg = {k: ck_cfg[k] for k in CFG_KEYS}
eng = ContinuousThoughtEngine(vocab_size=VOCAB_SIZE, **cfg)
eng.load_state_dict(sd, strict=True)
# The routing fix was applied once, at birth. Re-applying omega x4 would compound.
if os.environ.get("KURAMOTO_FIX_ON_RESUME", "0") == "1":
fix = apply_kuramoto_routing_fix(eng, log=log)
else:
fix = ck_cfg.get("kuramoto_fix", "applied at birth")
gate_temp = float(ck_cfg.get("gate_temp", os.environ.get("GATE_TEMP", "2.5")))
with torch.no_grad():
for blk in eng.blocks:
if hasattr(blk, "moe") and hasattr(blk.moe, "temperature"):
blk.moe.temperature = gate_temp
return eng, cfg, {"fresh": False, "resumed_from": str(ckpt_in), "kuramoto_fix": fix,
"parent": ck_cfg.get("parent", "hf:thefinalboss/fractus-cte")}
def resolve_start(ckpt_in: str | None) -> int:
raw = os.environ.get("START_TOKEN")
if not ckpt_in:
return int(raw or 0)
if raw is None:
raise SystemExit("resume: set START_TOKEN from the manifest (start_token_next)")
if int(raw) == 0:
raise SystemExit("resume: START_TOKEN=0 is legal only for a fresh model. Never rewind.")
return int(raw)
def routing_snapshot(eng) -> dict:
hits = getattr(eng, "_expert_hits", None)
out = {"lb": float(getattr(eng, "last_lb_loss", torch.tensor(float("nan"))))}
if hits is not None and hits.numel() and float(hits.sum()) > 0:
frac = hits / hits.sum()
out.update(alive=int((hits > 0).sum()), n=int(hits.numel()), max_frac=float(frac.max()))
return out
# ---------------------------------------------------------------- main
def main() -> None:
gpu = int(os.environ.get("GPU_ID", "0"))
n_gpu = max(1, int(os.environ.get("N_GPU", "1")))
scale = os.environ.get("SCALE", "1b")
if scale not in SCALES:
raise SystemExit(f"SCALE must be one of {sorted(SCALES)}")
lb_coef = float(os.environ.get("LB_COEF", "0.02"))
lr = float(os.environ.get("LR", "7e-4"))
opt_name = os.environ.get("OPT", "sgd")
ema_beta = float(os.environ.get("EMA_BETA", "0.98"))
ss_rate = float(os.environ.get("SS_RATE", "1.0"))
ss_p0 = float(os.environ.get("SS_PROB_START", "0.2"))
ss_p1 = float(os.environ.get("SS_PROB_END", "0.5"))
ss_ramp = int(os.environ.get("SS_RAMP_TOKENS", "50000000"))
repeat_coef = float(os.environ.get("REPEAT_COEF", "0.1"))
probe_every = int(os.environ.get("PROBE_EVERY", "200"))
log_every = int(os.environ.get("LOG_EVERY", "40"))
save_every = int(os.environ.get("SAVE_EVERY", "4000"))
max_steps = int(os.environ.get("MAX_STEPS", "0"))
B = int(os.environ.get("BATCH", "4"))
SEQ = int(os.environ.get("SEQ", "128"))
ce_chunk = int(os.environ.get("CE_CHUNK", "2048"))
block_ckpt = os.environ.get("BLOCK_CKPT", "0") == "1"
use_compile = os.environ.get("COMPILE", "0") == "1"
sync_timing = os.environ.get("SYNC_TIMING", "1") == "1"
corpus_path = os.environ.get("CORPUS", str(ROOT / "data" / "atom_corpus.i16"))
ckpt_in = os.environ.get("CKPT_IN") or None
out_dir = Path(os.environ.get("OUT_DIR", str(ROOT / "checkpoints" / "atom")))
ckpt_out = Path(os.environ.get("CKPT_OUT", str(out_dir / f"fractus_atom_{scale}_gpu{gpu}.pt")))
manifest_out = ckpt_out.parent / f"RESUME_MANIFEST_atom_gpu{gpu}.json"
log = lambda *a: print(f"GPU {gpu}:", *a, flush=True) # noqa: E731
if torch.cuda.is_available():
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
device = torch.device("cuda:0")
autocast = lambda: torch.autocast("cuda", dtype=torch.bfloat16) # noqa: E731
sync = torch.cuda.synchronize
else:
from contextlib import nullcontext
device = torch.device("cpu")
autocast = nullcontext
sync = lambda: None # noqa: E731
if not sync_timing:
sync = lambda: None # noqa: E731
torch.manual_seed(42 + gpu)
corpus, n_total = open_corpus(corpus_path)
lo = gpu * n_total // n_gpu
hi = (gpu + 1) * n_total // n_gpu
shard = corpus[lo:hi]
shard_len = hi - lo
step_ids = B * SEQ
# B lanes. Row b always reads lane b, so the carry of row b at step n+1
# continues exactly the text row b saw at step n. The old layout
# (one contiguous block viewed as (B, SEQ)) handed row b a carry from text
# (B-1)*SEQ ids earlier, not its own predecessor, whenever B > 1.
lane_len = shard_len // B
if lane_len < SEQ + 2:
raise SystemExit(f"corpus slice {lo}..{hi} too short for {B} lanes of {SEQ + 1} ids")
lane_starts = np.arange(B, dtype=np.int64) * lane_len
# START_TOKEN / start_token_next is the offset inside each lane.
start = resolve_start(ckpt_in)
start = (start // SEQ) * SEQ
eng, cfg, birth = build_engine(scale, ckpt_in, log)
n_params = sum(p.numel() for p in eng.parameters())
log(f"ATOM scale={scale} params={n_params:,} vocab={VOCAB_SIZE} B={B} SEQ={SEQ} "
f"attn={os.environ['FRACTUS_ATTN_IMPL']} block_ckpt={block_ckpt} ss_rate={ss_rate} "
f"opt={opt_name} lr={lr} fresh={birth['fresh']}")
log(f"corpus {corpus_path} total={n_total:,} slice={lo:,}..{hi:,} "
f"lanes={B}x{lane_len:,} lane_offset={start:,}")
eng = eng.to(device)
eng.reset_thought(B)
if use_compile:
# Compile the pure core of each block, not the engine. The engine
# compile was bypassed: every pass called payload(), which returned
# eng._orig_mod. The pure core has fixed shapes, so Inductor can fuse
# gelu, bias, scales, Kuramoto and LayerNorm. Works with BLOCK_CKPT
# because the checkpoint calls this same method.
try:
for blk in eng.blocks:
blk._tick_chunk_core_pure = torch.compile(
blk._tick_chunk_core_pure, dynamic=False)
log("torch.compile ON (_tick_chunk_core_pure per block)")
except Exception as exc: # pragma: no cover
log(f"compile skip: {exc}")
payload = lambda: eng # noqa: E731
if opt_name == "adamw":
opt = torch.optim.AdamW(eng.parameters(), lr=lr, weight_decay=0.01)
else:
opt = torch.optim.SGD(eng.parameters(), lr=lr, momentum=0.9)
def fetch(off: int) -> torch.Tensor:
"""(B, SEQ+1): row b = lane b, ids [off, off+SEQ]."""
rows = [np.asarray(shard[s + off:s + off + SEQ + 1], dtype=np.int64) for s in lane_starts]
return torch.from_numpy(np.stack(rows)).to(device, non_blocking=True)
def reset_rows(mask: torch.Tensor) -> None:
"""Zero the carry of the rows in mask only. Other rows keep their document."""
e = payload()
keep = (~mask).to(e.thought_state.dtype)
e.thought_state = e.thought_state * keep.view(-1, 1, 1)
for blk in e.blocks:
blk.attn_S = blk.attn_S * keep.view(-1, 1, 1, 1).to(blk.attn_S.dtype)
blk.attn_z = blk.attn_z * keep.view(-1, 1, 1).to(blk.attn_z.dtype)
blk.kuramoto_phases = blk.kuramoto_phases * keep.view(-1, 1, 1).to(blk.kuramoto_phases.dtype)
ckpt_out.parent.mkdir(parents=True, exist_ok=True)
def save(off: int) -> None:
ids_done = off * B
tmp = ckpt_out.with_suffix(ckpt_out.suffix + ".tmp")
torch.save({
"model_state": payload().state_dict(),
"config": {
**cfg, "vocab_size": VOCAB_SIZE, "atom": True, "scale": scale,
"gpu": gpu, "n_gpu": n_gpu, "corpus": str(corpus_path),
"corpus_slice": [lo, hi], "ids_processed": ids_done,
"layout": "lanes", "lane_len": lane_len,
"start_token_next": off, "batch": B, "seq": SEQ, "lr": lr, "opt": opt_name,
"gate_temp": float(os.environ.get("GATE_TEMP", "2.5")),
"kuramoto_fix": birth["kuramoto_fix"], "parent": birth["parent"],
"fresh_start_token": 0 if birth["fresh"] else None,
"trainer": "fast4gpu_atom",
},
}, tmp)
os.replace(tmp, ckpt_out)
mtmp = manifest_out.with_suffix(".json.tmp")
mtmp.write_text(json.dumps({
"ts": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
"trainer": "fast4gpu_atom", "gpu": gpu, "n_gpu": n_gpu, "scale": scale,
"vocab_size": VOCAB_SIZE, "ids_processed": ids_done,
"layout": "lanes", "batch": B, "lane_len": lane_len,
"start_token_next": off,
"note": "start_token_next is the offset inside each lane; resume with the same BATCH",
"corpus": str(corpus_path), "corpus_slice": [lo, hi], "ckpt": ckpt_out.name,
}, indent=1))
os.replace(mtmp, manifest_out)
log(f"saved -> {ckpt_out} @ lane offset {off:,} ({ids_done:,} ids)")
def probe(ids_done: int) -> None:
snap = snapshot_carry(payload())
p = unique40_probe(payload(), max_new=40, mode="prefix")
payload().reset_thought(B)
restore_carry(payload(), snap)
mean_u = sum(r["unique"] for r in p["rows"]) / max(1, len(p["rows"]))
log(f"UNIQUE@40 PREFIX mean_unique={mean_u:.1f} @ {ids_done:,} ids | routing {routing_snapshot(payload())}")
for r in p["rows"]:
log(f" {r['prompt']!r}: unique={r['unique']}/{r['n']} -> {r['text']!r}")
ema_tf = ema_ss = None
t_tf = t_ss = t_probe = 0.0
ids_tf = ids_ss = 0
n = 0
off = start
t_wall = time.time()
last_off = lane_len - SEQ - 1
while off <= last_off:
block = fetch(off)
chunk = block[:, :SEQ]
target = block[:, 1:]
off += SEQ
# Reset at each BOS so the model sees an empty thought at document
# boundaries. Continuity stays inside a document. Without this, a long
# run almost never sees an empty state, and generation (which starts
# from zero) is a regime the model never trained on. Measured: 81%
# accuracy with the training carry, 3% after reset, on a memorized phrase.
# Per row: a BOS in lane b resets lane b only, not the other documents.
bos_rows = (chunk == 257).any(dim=1)
if bool(bos_rows.any()):
reset_rows(bos_rows)
pos = off * B
# Snapshot only when the SS pass will run. Cloning the carry every
# step costs ~105 MB per row (16 blocks of 1280x1280), which is
# ~6.7 GB at B=64, for a restore that never happens when SS is off.
do_ss = ss_rate > 0 and should_ss(ss_rate)
carry = snapshot_carry(payload()) if do_ss else None
sync(); t0 = time.time()
with autocast():
loss, ex = v4_forward_losses(
payload(), chunk, target,
lb_coef=lb_coef, repeat_coef=repeat_coef,
ce_chunk=ce_chunk, block_ckpt=block_ckpt,
)
loss.backward()
torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
opt.step()
opt.zero_grad(set_to_none=True)
sync(); t_tf += time.time() - t0
ids_tf += step_ids
tf_v = float(ex["ce"].detach())
ema_tf = tf_v if ema_tf is None else ema_beta * ema_tf + (1 - ema_beta) * tf_v
ss_v = None
if do_ss:
restore_carry(payload(), carry)
sync(); t0 = time.time()
with autocast():
loss2, ce_ss = v4_ss_pass(
payload(), chunk, target, ex["h"].detach(),
ss_prob=ss_prob, lb_coef=lb_coef, ce_chunk=ce_chunk, block_ckpt=False,
)
loss2.backward()
torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
opt.step()
opt.zero_grad(set_to_none=True)
sync(); t_ss += time.time() - t0
ids_ss += step_ids
ss_v = float(ce_ss.detach())
ema_ss = ss_v if ema_ss is None else ema_beta * ema_ss + (1 - ema_beta) * ss_v
n += 1
if n % log_every == 0:
wall = max(time.time() - t_wall, 1e-9)
mem = ""
if device.type == "cuda":
mem = f" mem_peak={torch.cuda.max_memory_allocated() / 1e9:.1f}GB"
ss_s = f" ss={ss_v:.3f} ema_ss={ema_ss:.3f}" if ss_v is not None else ""
log(
f"{pos:>12,} tf={tf_v:.3f} ema_tf={ema_tf:.3f}{ss_s} "
f"rep={float(ex['repeat'].detach()):.3f} lb={float(ex['lb'].detach()):.3f} | "
f"data {ids_tf / wall:.0f} ids/s wall | "
f"TF {ids_tf / max(t_tf, 1e-9):.0f} ids/s ({t_tf / wall:.0%}) "
f"SS {ids_ss / max(t_ss, 1e-9):.0f} ids/s ({t_ss / wall:.0%}) "
f"probe {t_probe / wall:.0%}{mem}"
)
if probe_every > 0 and n % probe_every == 0:
t0 = time.time(); probe(pos); t_probe += time.time() - t0
if save_every > 0 and n % save_every == (gpu * 500) % save_every:
save(off)
if max_steps and n >= max_steps:
break
wall = max(time.time() - t_wall, 1e-9)
save(off)
summary = {
"steps": n, "ids_tf": ids_tf, "ids_ss": ids_ss, "wall_s": round(wall, 2),
"data_ids_per_s_wall": round(ids_tf / wall, 1),
"tf_ids_per_s": round(ids_tf / max(t_tf, 1e-9), 1),
"ss_ids_per_s": round(ids_ss / max(t_ss, 1e-9), 1) if ids_ss else None,
"share": {"tf": round(t_tf / wall, 3), "ss": round(t_ss / wall, 3),
"probe": round(t_probe / wall, 3)},
"ema_tf": ema_tf, "params": n_params, "scale": scale, "batch": B, "seq": SEQ,
"block_ckpt": block_ckpt, "compile": use_compile, "device": str(device),
}
if device.type == "cuda":
summary["mem_peak_gb"] = round(torch.cuda.max_memory_allocated() / 1e9, 2)
log("SUMMARY " + json.dumps(summary))
if __name__ == "__main__":
main()