spec100m / code /bench_cascade.py
Akahsizrr's picture
Upload code/bench_cascade.py with huggingface_hub
af8fc3a verified
Raw History Blame Contribute Delete
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
@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()