File size: 5,674 Bytes
af8fc3a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
"""Cascade benchmark: verified vs optimistic composite decode.

Usage: python3 bench_cascade.py <ckpt> [n_tokens] [batch]
"""
import sys
import time
import torch
import tiktoken

from config import Config
from model import build_model
from inference import InferenceEngine
from cascade import CascadeEngine
from sam import build_from_ids


@torch.no_grad()
def stream_ppl(model, ids, chunk=1024):
    """Base-model perplexity of an emitted stream — honest quality gate.
    Degenerate loops score high PPL; coherent text scores low."""
    nll, cnt = 0.0, 0
    for i in range(0, len(ids) - 1, chunk):
        seg = ids[i:i + chunk + 1]
        x = torch.tensor([seg[:-1]], device="cuda")
        h = model(x)
        lg = model.lm_head(h).float()
        tgt = torch.tensor(seg[1:], device="cuda")
        lp = torch.log_softmax(lg[0], -1)
        nll -= lp.gather(-1, tgt.unsqueeze(-1)).sum().item()
        cnt += len(tgt)
    import math
    return math.exp(nll / max(1, cnt))


def main():
    ckpt_path = sys.argv[1] if len(sys.argv) > 1 else \
        "/root/checkpoints/spec_v7/best.pt"
    n_tokens = int(sys.argv[2]) if len(sys.argv) > 2 else 256
    batch = int(sys.argv[3]) if len(sys.argv) > 3 else 1

    cfg = Config.v5_500m()
    cfg.max_seq_len = min(cfg.max_seq_len, 1 << 20)
    model = build_model(cfg, "cuda")
    ckpt = torch.load(ckpt_path, map_location="cuda")
    model.load_state_dict(ckpt["model"], strict=False)
    # only enable markov drafting if the ckpt actually has trained params —
    # injecting random markov weights sandbags v6 checkpoints
    cfg.use_markov_head = any("spec_markov_emb" in k for k in ckpt["model"])
    model.eval()
    print(f"loaded {ckpt_path} (step {ckpt.get('step')}, "
          f"loss {ckpt.get('loss'):.4f})  batch={batch}")

    enc = tiktoken.get_encoding("gpt2")
    prompts = [
        "The quick brown fox jumps over the lazy dog. In a surprising turn of events,",
        "The history of the Roman Empire begins in 753 BC with the founding of Rome.",
        "To make a simple loaf of bread, you need flour, water, salt, and yeast.",
        "Photosynthesis is the process by which plants convert sunlight into energy.",
    ]
    pids = [enc.encode_ordinary(p) for p in prompts]

    # SAM corpus: a slice of the cached mixture (self-consistent text)
    sam = None
    try:
        from data_pipeline import load_tokens
        toks = load_tokens("mixture500m")
        sam = build_from_ids(toks[:2_000_000].tolist())
        print(f"SAM corpus: {len(sam.text):,} tokens, {len(sam.length):,} states")
    except Exception as e:
        print(f"SAM corpus unavailable ({e}) — heads-only mode")

    eng = InferenceEngine(model, "cuda")
    eng.alloc_cache(batch=batch, max_len=1 << 16)
    cas = CascadeEngine(model, eng, sam=sam)
    # warmup: trigger draft-graph capture + lazy init OUTSIDE timed region
    cas.generate(pids[0], 64, mode="optimistic", sam_extend=0, batch=batch)

    for pi, prompt in enumerate(prompts[:2]):
        print(f"\n=== {prompt[:60]!r} ===")
        for mode, sam_ext, samp in [("verified", 0, False),
                                    ("optimistic", 0, False),
                                    ("optimistic", 512, False),
                                    ("optimistic", 0, True)]:
            outs, stats = cas.generate(
                pids[pi] if batch == 1 else pids, n_tokens,
                mode=mode, sam_extend=sam_ext, batch=batch,
                sample=samp, temperature=0.8, top_p=0.9)
            print(f"  {mode:>10} sam={sam_ext:4d} samp={int(samp)}: "
                  f"{stats['tok_s']:>10,.0f} tok/s aggregate "
                  f"({stats['per_stream_tok_s']:,.0f}/stream)  "
                  f"commit/round {stats['committed_per_round']:.1f}  "
                  f"soft-accept {stats['soft_accept']:.2f}  "
                  f"exact {stats['exact_accept']:.2f}  "
                  f"sam {stats['sam_tokens']}")
        print(f"  sample: {enc.decode(outs[0][:80])!r}")
        print(f"  stream base-PPL: {stream_ppl(model, pids[pi] + outs[0]):.1f}")

    # batch scaling sweep — small per-stream cache so big B fits
    if batch == 1:
        print("\n=== batch scaling (optimistic, sam=256) ===")
        for b in [1, 4, 16, 64, 128, 256, 512, 1024]:
            try:
                eng.alloc_cache(batch=b, max_len=8192)
                cas.eng = eng
                cas.generate(pids, 64, mode="optimistic", sam_extend=0,
                             batch=b)                    # warmup + capture
                outs, st = cas.generate(pids, n_tokens, mode="optimistic",
                                        sam_extend=256, batch=b)
                print(f"  B={b:3d}: {st['tok_s']:>10,.0f} tok/s aggregate "
                      f"({st['per_stream_tok_s']:,.0f}/stream)  "
                      f"soft-accept {st['soft_accept']:.2f}  "
                      f"sam {st['sam_tokens']}")
            except torch.cuda.OutOfMemoryError:
                print(f"  B={b:3d}: OOM"); break

        print("\n=== sam_extend sweep (optimistic, B=1) ===")
        eng.alloc_cache(batch=1, max_len=1 << 16)
        for se in [0, 512, 2048, 8192, 32768]:
            outs, st = cas.generate(pids[0], n_tokens * 4,
                                    mode="optimistic", sam_extend=se,
                                    batch=1)
            print(f"  sam={se:6d}: {st['tok_s']:>10,.0f} tok/s  "
                  f"sam-tokens {st['sam_tokens']}  "
                  f"soft-accept {st['soft_accept']:.2f}")
            print(f"       text: {enc.decode(outs[0][:60])!r}")

    print("BENCH_DONE")


if __name__ == "__main__":
    main()