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())