"""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