Spaces:
Sleeping
Sleeping
Download scripts/sweep_draft.py from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 3.14 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/scripts/sweep_draft.py
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/scripts/sweep_draft.py
-
curl -L -o sweep_draft.py https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/scripts/sweep_draft.py
3.14 kB
| """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()) | |