spec100m / code /cascade.py
Akahsizrr's picture
Upload code/cascade.py with huggingface_hub
c7c7c78 verified
Raw History Blame Contribute Delete
11.8 kB
"""v8 Cascade Engine — composite decode for extreme throughput.
Core trick: the free-run forward IS the verify pass. Forwarding a drafted
block with causal_extend=True gives every drafted position its causal
hidden state for free — scoring argmax(lm_head(h_i)) vs the draft verifies
the whole block inside the same forward that advances the KV cache.
Stream tiers per round:
- heads draft K tokens (v7 markov-conditioned, confidence head)
- optional SAM suffix-automaton extension emitted WITHOUT forwarding
(the honest "composite" component of the throughput number)
Modes:
"verified" — commit longest prefix where draft == base argmax
(or p_base >= tau for soft acceptance). Batch-aligned via
min-commit: every stream advances by the worst stream's
accept length. Output = base-consistent.
"optimistic" — commit the WHOLE drafted block + SAM extension; the
soft_accept stat reports how much of it the base would
have kept. The 250K+ path.
Batch: B>1 streams share every forward — a 271M model at B=32 costs ~the
same as B=1, multiplying aggregate tok/s.
"""
from __future__ import annotations
import time
import torch
from sam import SuffixAutomaton
class CascadeEngine:
def __init__(self, model, engine, device="cuda",
sam: SuffixAutomaton = None, tau: float = 0.35):
self.model = model
self.eng = engine
self.cfg = model.cfg
self.device = device
self.tau = tau
self.sam = sam
self._dg = None # captured draft graph
self._dg_h = None
self._dg_hist = None
self._dg_out = None
self._dg_conf = None
# ------------------------------------------------------------ draft graph
def capture_draft(self, batch: int):
"""Capture spec_draft as a CUDA graph — the 32-group sequential
loop is launch-bound in eager; graphed it replays in ~1-3ms."""
cfg = self.cfg
B = batch
self._dg_h = torch.zeros(B, cfg.d_model, device=self.device,
dtype=self.eng.dtype)
self._dg_hist = torch.zeros(B, cfg.medusa_cond_group,
dtype=torch.long, device=self.device)
for _ in range(3):
self.model.spec_draft(self._dg_h, self._dg_hist, return_conf=True)
torch.cuda.synchronize()
self._dg = torch.cuda.CUDAGraph()
with torch.cuda.graph(self._dg):
self._dg_out, self._dg_conf = self.model.spec_draft(
self._dg_h, self._dg_hist, return_conf=True)
torch.cuda.synchronize()
def _draft_graphed(self, h_anchor, hist):
if self._dg is None or self._dg_h.shape[0] != h_anchor.shape[0]:
self.capture_draft(h_anchor.shape[0])
self._dg_h.copy_(h_anchor)
self._dg_hist.copy_(hist)
self._dg.replay()
return self._dg_out.clone(), self._dg_conf
# ------------------------------------------------------------ scoring
@torch.no_grad()
def _score(self, h_all, h_anchor, block):
"""h_all [B,L,d] block hiddens (causal), h_anchor [B,d] pre-block.
Returns (exact [B,L], soft [B,L], base_argmax [B,L]).
Chunked over L so [B,L,V] logits/probs never materialize fully."""
B, L = block.shape
prev_h = torch.cat([h_anchor.unsqueeze(1), h_all[:, :-1]], dim=1)
exact = torch.empty(B, L, dtype=torch.bool, device=self.device)
soft = torch.empty(B, L, dtype=torch.bool, device=self.device)
argmax = torch.empty(B, L, dtype=torch.long, device=self.device)
CH = 64
for i in range(0, L, CH):
lg = self.model.lm_head(prev_h[:, i:i + CH]) # [B,ch,V]
am = lg.argmax(-1)
argmax[:, i:i + CH] = am
ex = (block[:, i:i + CH] == am)
exact[:, i:i + CH] = ex
# p_base(draft_tok) without materializing the full softmax:
# p = exp(logit_tok - logsumexp(logits))
lse = lg.float().logsumexp(-1) # [B,ch]
lt = lg.gather(-1, block[:, i:i + CH]
.unsqueeze(-1)).squeeze(-1).float()
soft[:, i:i + CH] = ex | ((lt - lse).exp() >= self.tau)
return exact, soft, argmax
# ------------------------------------------------------------ generate
@torch.no_grad()
def generate(self, prompt_ids, n_tokens, mode="verified",
sam_extend=0, conf_gate=0.0, batch=1, use_markov=None,
sample=False, temperature=1.0, top_p=0.9):
"""prompt_ids: list[int] (shared) or list[list[int]] (per-stream).
sample=True draws drafted tokens from head distributions (temp/top-p)
instead of argmax — non-degenerate text for optimistic mode.
Returns (streams, stats)."""
cfg = self.cfg
G, K = cfg.medusa_cond_group, cfg.medusa_heads
B = batch
if use_markov is not None:
cfg.use_markov_head = use_markov
prompts = ([prompt_ids] * B if isinstance(prompt_ids[0], int)
else [prompt_ids[i % len(prompt_ids)] for i in range(B)])
P = max(len(p) for p in prompts)
self.eng.reset_cache()
ids = torch.zeros(B, P, dtype=torch.long, device=self.device)
for i, p in enumerate(prompts):
ids[i, :len(p)] = torch.tensor(p, device=self.device)
h_all0 = self.eng.step(ids, start_pos=0)
pos = P
h_anchor = h_all0[:, -1] # [B,d]
committed = [list(p) for p in prompts]
outs = [[] for _ in range(B)]
pending = None # [B,1] or None
t0 = time.perf_counter()
n_rounds = n_commit = n_soft = n_exact = n_draft = n_sam = 0
while min(len(o) for o in outs) < n_tokens:
# ---- heads draft ----
hist = torch.stack([
torch.tensor(
([committed[b][0]] * max(0, G - len(committed[b]))
+ committed[b][-G:])[-G:], device=self.device)
for b in range(B)])
if mode == "optimistic" and not sample:
draft, conf = self._draft_graphed(h_anchor, hist)
else:
draft = self.model.spec_draft(
h_anchor, hist,
first_head=0 if pending is None else 1,
prefix_ids=pending,
sample=sample, temperature=temperature,
top_p=top_p) # [B,K-first]
if pending is not None:
block = torch.cat([pending, draft], dim=1) # [B,L]
else:
block = draft # [B,L]
L = block.shape[1]
# ---- one forward: advances cache AND verifies ----
h_all = self.model(block, self.eng._cache_k, self.eng._cache_v,
pos, causal_extend=True) # [B,L,d]
exact, soft, base_am = self._score(h_all, h_anchor, block)
# vectorized accept lengths (one sync, no per-stream nonzero)
if mode == "verified":
first = 1 if pending is not None else 0
if first:
# pending is a base-sampled/argmax commit — it is
# unconditionally accepted (the base chose it);
# a low-prob sample must not fail verification
soft = soft.clone()
soft[:, 0] = True
failmask = ~soft
# first fail index per stream, or L if none
has_fail = failmask.any(1)
j_b = torch.where(
has_fail,
failmask.int().argmax(1),
torch.full((B,), L, device=self.device))
j = int(j_b.min().item())
soft_c = soft[:, :j].sum().item()
exact_c = exact[:, :j].sum().item()
for b in range(B):
take = block[b, first:j].tolist()
outs[b].extend(take)
committed[b].extend(take)
n_soft += soft_c
n_exact += exact_c
n_draft += j * B
if j < L:
if sample:
# correction sampled from the BASE distribution —
# keeps the stream coherent instead of lock-step
# greedy. Approximate-rejection-consistent: the
# accepted prefix passed p>=tau under the base.
prev_h = torch.cat([h_anchor.unsqueeze(1),
h_all[:, :-1]], dim=1)
lg = self.model.lm_head(prev_h[:, j]).float() \
/ max(temperature, 1e-4) # [B,V]
if top_p < 1.0:
s, si = lg.sort(-1, descending=True)
sp = s.softmax(-1)
rm = sp.cumsum(-1) - sp >= top_p
lg = si.gather(-1, torch.multinomial(
s.masked_fill(rm, float("-inf"))
.softmax(-1), 1))
else:
lg = torch.multinomial(lg.softmax(-1), 1)
pending = lg
else:
pending = base_am[:, j:j + 1]
else:
pending = self.model.lm_head(h_all[:, -1]).argmax(-1)
pending = pending.unsqueeze(1)
pl = pending[:, 0].tolist()
for b in range(B):
outs[b].append(pl[b])
committed[b].append(pl[b])
n_commit += (j - first + 1) * B
h_anchor = h_all[:, j - 1] if j > 0 else h_anchor
pos += j
n_rounds += 1
else: # optimistic: whole block commits
n_soft += int(soft.sum().item())
n_exact += int(exact.sum().item())
block_l = block.tolist()
for b in range(B):
take = block_l[b]
outs[b].extend(take)
committed[b].extend(take)
if self.sam is not None and sam_extend > 0:
# index the committed head-tokens too (self-consistent
# pool) then draft the continuation
self.sam.extend_many(take)
ext = self.sam.draft(committed[b][-128:], sam_extend,
min_len=2)
if ext:
outs[b].extend(ext)
committed[b].extend(ext)
self.sam.extend_many(ext)
n_sam += len(ext)
n_draft += L * B
n_commit += L * B
h_anchor = h_all[:, -1]
pending = None
pos += L
n_rounds += 1
torch.cuda.synchronize()
dt = time.perf_counter() - t0
n_out = min(len(o) for o in outs)
stats = {
"tok_s": n_out * B / dt,
"per_stream_tok_s": n_out / dt,
"rounds": n_rounds,
"committed_per_round": n_commit / max(1, n_rounds * B),
"soft_accept": n_soft / max(1, n_draft),
"exact_accept": n_exact / max(1, n_draft),
"sam_tokens": n_sam,
"batch": B, "mode": mode,
}
return [o[:n_tokens] for o in outs], stats