"""Cascade benchmark: verified vs optimistic composite decode. Usage: python3 bench_cascade.py [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()