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