Spaces:
Sleeping
Sleeping
File size: 3,140 Bytes
421b8c2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 | """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())
|