Download code/cascade.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/cascade.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/cascade.py
-
curl -L -o cascade.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/cascade.py
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 | |
| 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 | |
| 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 | |