"""Sweep neural-draft skip sets x K on Surjo-50m: accept rate + tok/s vs K0. Uses the native session directly (Engine gives no access to set_draft_steps). Run from repo root: .venv/Scripts/python scripts/sweep_draft.py """ from __future__ import annotations import time from cism import _native, loader from cism.engine import Engine SURJO = r"G:\hf_cache\models--SurjoLabs--Surjo-50m\snapshots\35c9fa6eacd0e8e977c3e7d878a771837164f930" PROMPT = "The future of artificial intelligence is" PRECISION = "fp32" THREADS = 2 MAX_NEW = 128 REPS = 3 # Plan (verified): 0 L0 XSA, 1-3 GDN p0, 4 XSA, 5-7 GDN p0, 8 XSA, # 9-11 GDN p1, 12 XSA p1, 13-15 GDN p1, 16 XSA p1, 17 XSA coda. ALL_GDN = [1, 2, 3, 5, 6, 7, 9, 10, 11, 13, 14, 15] P1_GDN = [9, 10, 11, 13, 14, 15] ALL_XSA_MID = [4, 8, 12, 16] EVERYTHING = list(range(18)) SKIP_SETS: dict[str, list[int] | None] = { "K0 (no spec)": None, "default p1-GDN": P1_GDN, "all GDN": ALL_GDN, "all GDN + mid XSA": ALL_GDN + ALL_XSA_MID, "everything": EVERYTHING, } KS = (2, 4) def median(xs: list[float]) -> float: return sorted(xs)[len(xs) // 2] def main() -> int: loaded = loader.import_model(SURJO) enc = Engine.from_weights(loaded.config, loaded.weights, loaded.tokenizer) prompt_ids = enc.encode(PROMPT) print(f"prompt ids: {len(prompt_ids)} tokens") model = _native.SurjoModel(loaded.config, loaded.weights, PRECISION, THREADS, "fp32") # Ceiling probe: block-per-token (nll, one fused block call) vs sequential. seq_ids = prompt_ids + [1] * 64 t0 = time.perf_counter() model.nll(seq_ids) nll_dt = time.perf_counter() - t0 print(f"nll block: {(len(seq_ids) - 1) / nll_dt:.1f} tok/s " f"({nll_dt * 1000 / (len(seq_ids) - 1):.3f} ms/tok)") base_ids: list[int] | None = None for name, skips in SKIP_SETS.items(): ks = (0,) if skips is None else KS for k in ks: rates, toks, accs = [], [], [] for _ in range(REPS): s = model.create_session( list(prompt_ids), max_new_tokens=MAX_NEW, temperature=0.0, top_p=1.0, top_k=0, seed=0, eos_token_ids=[]) if skips is not None: s.set_draft_steps(skips) t0 = time.perf_counter() out = s.next_tokens(MAX_NEW, k) dt = time.perf_counter() - t0 rates.append(len(out) / dt) toks.append(out) p, a = s.spec_stats accs.append(a / p if p else -1.0) if base_ids is None: base_ids = toks[0] same = True else: same = all(t == base_ids for t in toks) print(f"{name:18s} K={k} {median(rates):7.1f} tok/s " f"accept={median(accs):5.1%} identical={same} " f"ntok={len(toks[0])}", flush=True) if not same: print(" TOKEN MISMATCH vs K0!") return 1 return 0 if __name__ == "__main__": raise SystemExit(main())