Download code/bench_cascade.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 5.67 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/bench_cascade.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/bench_cascade.py
-
curl -L -o bench_cascade.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/bench_cascade.py
5.67 kB
| """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 | |
| 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() | |