"""Parallel decoding (survey Def. 3): draft several future tokens, verify them in one pass. - Speculative decoding (Leviathan et al. 2023; Chen et al. 2023): the helper drafts. - Prompt lookup decoding: model-free drafts copied from an earlier n-gram match. - Jacobi decoding (Santilli et al. 2023): a block of guesses refined until it stops changing. """ from __future__ import annotations import torch from decoding.common import ( KV, ROUND, ar_decode, first_divergence, forward_stats, lcp, make_generator, make_run, now, sample, topk_entries, ) from prompting import prompt_of PQ_TOPK = 5 def predicted_speedup(alpha: float, gamma: int, c: float) -> float: """Expected wall-clock improvement of speculative decoding (Leviathan et al. 2023, Thm 3.8).""" if gamma <= 0: return 1.0 alpha = min(max(alpha, 0.0), 1.0) tokens = gamma + 1 if alpha >= 1.0 else (1 - alpha ** (gamma + 1)) / (1 - alpha) return tokens / (gamma * c + 1) def expected_tokens(alpha: float, gamma: int) -> float: alpha = min(max(alpha, 0.0), 1.0) return gamma + 1 if alpha >= 1.0 else (1 - alpha ** (gamma + 1)) / (1 - alpha) def verify(ps: torch.Tensor, qs: list[torch.Tensor], drafts: list[int], greedy: bool, gen): """Check drafts against target distributions ``ps`` [g+1, V]. Returns (n_accepted, next_token, kind, acceptance_probs). ``kind`` is "fix" when a draft was rejected and replaced, "bonus" when every draft passed. """ acc_probs: list[float] = [] for i, d in enumerate(drafts): if greedy: top = int(ps[i].argmax()) acc_probs.append(1.0 if top == d else 0.0) if top != d: return i, top, "fix", acc_probs else: ratio = min(1.0, float(ps[i][d]) / max(float(qs[i][d]), 1e-20)) acc_probs.append(ratio) if float(torch.rand((), generator=gen)) >= ratio: residual = (ps[i] - qs[i]).clamp_min(0) total = float(residual.sum()) residual = residual / total if total > 1e-12 else ps[i] return i, sample(residual, gen), "fix", acc_probs g = len(drafts) nxt = int(ps[g].argmax()) if greedy else sample(ps[g], gen) return g, nxt, "bonus", acc_probs def _truncate_at_stop(new: list[int], stop_ids: set[int]) -> tuple[list[int], bool]: for j, t in enumerate(new): if t in stop_ids: return new[: j + 1], True return new, False def _finish_method_run(fam, label, out, kinds, tok_iter, iters, finish, wall, **kvs) -> dict: run = make_run(label, fam.tok, out, finish, kinds=kinds, tok_iter=tok_iter, iters=iters, **forward_stats(**kvs)) run["time_ms"]["wall"] = round(1000 * wall, 2) return run def _baseline(fam, prompt_ids, P, greedy=True): return ar_decode( fam.main, fam.tok, prompt_ids, P["max_new_tokens"], fam.stop_ids, disp=fam.disp, greedy=greedy, temperature=P.get("temperature", 1.0), gen=make_generator(P.get("seed", 0)), label=f"Autoregressive ({fam.main_id.split('/')[-1]})", record=False, ) def _speed_summary(base: dict, run: dict, target_forwards: int) -> dict: emitted = len(run["ids"]) return { "tokens": emitted, "target_forwards": target_forwards, "baseline_forwards": base["n_forward"]["main"], "tau": round(emitted / max(target_forwards, 1), 3), "speedup_wall": round(base["time_ms"]["wall"] / max(run["time_ms"]["wall"], 1e-6), 3), "baseline_ms_per_token": round(base["time_ms"]["wall"] / max(len(base["ids"]), 1), 2), "method_ms_per_token": round(run["time_ms"]["wall"] / max(emitted, 1), 2), } def _identity(base: dict, run: dict) -> dict: n = min(len(base["ids"]), len(run["ids"])) div = first_divergence(base["ids"][:n], run["ids"][:n]) return {"identical": div is None, "first_divergence": div} # --------------------------------------------------------------------------- # Speculative decoding # --------------------------------------------------------------------------- @torch.no_grad() def run_speculative(fam, P: dict) -> dict: if fam.helper is None: raise ValueError(f"{fam.label} has no helper model, so speculative decoding is unavailable.") prompt_ids = prompt_of(fam, P) greedy = P["mode"] == "greedy" temp = P["temperature"] gamma, max_new = P["gamma"], P["max_new_tokens"] base = _baseline(fam, prompt_ids, P, greedy=greedy) gen = make_generator(P["seed"]) T, D = KV(fam.main), KV(fam.helper) S = list(prompt_ids) out: list[int] = [] kinds: list[str] = [] tok_iter: list[int] = [] iters: list[dict] = [] betas: list[float] = [] finish = "length" t0 = now(T.device) while len(out) < max_new: budget = min(gamma, max_new - len(out) - 1) drafts: list[int] = [] qs: list[torch.Tensor] = [] d_before = len(D.log) for _ in range(budget): z = D.run(S + drafts).logits[0, -1].float() q = torch.softmax(z / max(temp, 1e-5), dim=-1) if not greedy else torch.softmax(z, dim=-1) d = int(q.argmax()) if greedy else sample(q, gen) drafts.append(d) qs.append(q) if d in fam.stop_ids: break g = len(drafts) zt = T.run(S + drafts, keep=g + 1).logits[0].float() ps = torch.softmax(zt / max(temp, 1e-5), dim=-1) if not greedy else torch.softmax(zt, dim=-1) n, nxt, kind, acc_probs = verify(ps, qs, drafts, greedy, gen) beta = [round(float(torch.minimum(ps[i], qs[i]).sum()), ROUND) for i in range(min(n + 1, g))] betas.extend(beta) rows = [] for i, d in enumerate(drafts): status = "acc" if i < n else ("rej" if i == n else "unv") rows.append([d, fam.disp(d), round(float(qs[i][d]), ROUND), round(float(ps[i][d]), ROUND), round(acc_probs[i], ROUND) if i < len(acc_probs) else None, status]) new, stopped = _truncate_at_stop(drafts[:n] + [nxt], fam.stop_ids) emit_kinds = (["acc"] * n + [kind])[: len(new)] iters.append({ "k": len(iters), "pos": len(out), "draft": rows, "emit": [[t, fam.disp(t), k] for t, k in zip(new, emit_kinds)], "beta": beta, "t_draft_ms": round(1000 * sum(s for _, s in D.log[d_before:]), 3), "t_verify_ms": round(1000 * T.log[-1][1], 3), "top": { "p": [topk_entries(ps[i], fam.disp, PQ_TOPK) for i in range(g + 1)], "q": [topk_entries(qs[i], fam.disp, PQ_TOPK) for i in range(g)], }, }) for t, k in zip(new, emit_kinds): out.append(t) kinds.append(k) tok_iter.append(len(iters) - 1) S += new if stopped: finish = "eos" break wall = now(T.device) - t0 run = _finish_method_run(fam, f"Speculative ({'greedy' if greedy else 'sampling'}, γ={gamma})", out, kinds, tok_iter, iters, finish, wall, main=T, helper=D) drafted = sum(len(it["draft"]) for it in iters) accepted = sum(k == "acc" for k in kinds) alpha_hat = sum(betas) / len(betas) if betas else 0.0 c_hat = D.decode_forward_mean() / max(base["time_ms"]["main_fwd_mean"] / 1000, 1e-9) summary = { **_speed_summary(base, run, T.n_forward), "drafted": drafted, "accepted": accepted, "accept_rate": round(accepted / drafted, 3) if drafted else None, "alpha_hat": round(alpha_hat, 3), "c_hat": round(c_hat, 3), "gamma": gamma, "predicted_speedup": round(predicted_speedup(alpha_hat, gamma, c_hat), 3), "expected_tokens": round(expected_tokens(alpha_hat, gamma), 3), "curve": [[g, round(predicted_speedup(alpha_hat, g, c_hat), 3)] for g in range(1, 11)], } summary["gamma_star"] = max(summary["curve"], key=lambda r: r[1])[0] if greedy: summary.update(_identity(base, run)) return {"runs": {"baseline": base, "method": run}, "summary": summary} # --------------------------------------------------------------------------- # Prompt lookup decoding # --------------------------------------------------------------------------- def ngram_lookup(S: list[int], n_max: int, n_min: int, k: int) -> tuple[list[int], tuple[int, int] | None]: """Draft tokens copied from an earlier occurrence of the longest matching tail n-gram. Among occurrences of that n-gram, prefer the longest continuation (up to k tokens), then the most recent one. Matches right before the end are cut short by the end of the sequence, which matters for self-overlapping patterns. """ if k <= 0: return [], None for n in range(min(n_max, len(S) - 1), n_min - 1, -1): tail = S[-n:] best: tuple[list[int], int] | None = None for st in range(len(S) - n - 1, -1, -1): if S[st : st + n] == tail: cont = S[st + n : st + n + k] if cont and (best is None or len(cont) > len(best[0])): best = (cont, st + n) if len(cont) == k: break if best is not None: return best[0], (best[1], n) return [], None @torch.no_grad() def run_pld(fam, P: dict) -> dict: prompt_ids = prompt_of(fam, P) max_new, k = P["max_new_tokens"], P["num_pred"] base = _baseline(fam, prompt_ids, P) T = KV(fam.main) S = list(prompt_ids) n_prompt = len(S) out: list[int] = [] kinds: list[str] = [] tok_iter: list[int] = [] iters: list[dict] = [] finish = "length" t0 = now(T.device) while len(out) < max_new: t_look = now(T.device) cands, src = ngram_lookup(S, P["ngram_max"], P["ngram_min"], min(k, max_new - len(out) - 1)) t_look = now(T.device) - t_look g = len(cands) preds = T.run(S + cands, keep=g + 1).logits[0].argmax(-1).tolist() n = lcp(cands, preds[:g]) kind = "ar" if g == 0 else ("bonus" if n == g else "fix") new, stopped = _truncate_at_stop(cands[:n] + [preds[n]], fam.stop_ids) emit_kinds = (["acc"] * n + [kind])[: len(new)] rows = [[d, fam.disp(d), None, None, 1.0 if i < n else 0.0, "acc" if i < n else ("rej" if i == n else "unv")] for i, d in enumerate(cands)] iters.append({ "k": len(iters), "pos": len(out), "draft": rows, "emit": [[t, fam.disp(t), kk] for t, kk in zip(new, emit_kinds)], "t_draft_ms": round(1000 * t_look, 3), "t_verify_ms": round(1000 * T.log[-1][1], 3), "pld": None if src is None else { "start": src[0], "n": src[1], "in_prompt": src[0] < n_prompt, "ngram": [fam.disp(t) for t in S[src[0] - src[1] : src[0]]], }, }) for t, kk in zip(new, emit_kinds): out.append(t) kinds.append(kk) tok_iter.append(len(iters) - 1) S += new if stopped: finish = "eos" break wall = now(T.device) - t0 run = _finish_method_run(fam, f"Prompt lookup (n≤{P['ngram_max']}, k={k})", out, kinds, tok_iter, iters, finish, wall, main=T) drafted = sum(len(it["draft"]) for it in iters) accepted = sum(kk == "acc" for kk in kinds) summary = { **_speed_summary(base, run, T.n_forward), "drafted": drafted, "accepted": accepted, "accept_rate": round(accepted / drafted, 3) if drafted else None, "draft_hits": sum(1 for it in iters if it["draft"]), **_identity(base, run), } return {"runs": {"baseline": base, "method": run}, "summary": summary} # --------------------------------------------------------------------------- # Jacobi decoding # --------------------------------------------------------------------------- def _init_fill(S: list[int], m: int, how: str, gen, vocab: int) -> list[int]: if how == "random": return torch.randint(0, vocab, (m,), generator=gen).tolist() return [S[-1]] * m # "repeat": repeat the last known token @torch.no_grad() def run_jacobi(fam, P: dict) -> dict: prompt_ids = prompt_of(fam, P) max_new, m = P["max_new_tokens"], P["block"] base = _baseline(fam, prompt_ids, P) gen = make_generator(P.get("seed", 0)) vocab = min(len(fam.tok), fam.vocab_size) T = KV(fam.main) S = list(prompt_ids) out: list[int] = [] kinds: list[str] = [] tok_iter: list[int] = [] iters: list[dict] = [] finish = "length" y = _init_fill(S, m, P["init"], gen, vocab) t0 = now(T.device) while len(out) < max_new: # at most g + 1 tokens are committed, so g <= remaining - 1 keeps us within budget y = y[: min(m, max_new - len(out) - 1)] g = len(y) yn = T.run(S + y, keep=g + 1).logits[0].argmax(-1).tolist() a = lcp(y, yn[:g]) # leading guesses that reproduced themselves are exact commit = y + [yn[g]] if a == g else yn[: a + 1] new, stopped = _truncate_at_stop(commit, fam.stop_ids) last_kind = "ar" if g == 0 else ("bonus" if a == g else "fix") emit_kinds = ["acc"] * (len(new) - 1) + [last_kind] iters.append({ "k": len(iters), "pos": len(out), "draft": [[d, fam.disp(d), None, None, 1.0 if i < a else 0.0, "acc" if i < a else ("rej" if i == a else "unv")] for i, d in enumerate(y)], "emit": [[t, fam.disp(t), kk] for t, kk in zip(new, emit_kinds)], "t_verify_ms": round(1000 * T.log[-1][1], 3), "jacobi": { "before": [[t, fam.disp(t)] for t in y], "after": [[t, fam.disp(t)] for t in yn[:g]], "fixed": len(new), }, }) for t, kk in zip(new, emit_kinds): out.append(t) kinds.append(kk) tok_iter.append(len(iters) - 1) S += new if stopped: finish = "eos" break # carry the unconverged guesses forward and refill the block y = (yn[a + 1 : g] + _init_fill(S, m, P["init"], gen, vocab))[:m] wall = now(T.device) - t0 run = _finish_method_run(fam, f"Jacobi (block m={m})", out, kinds, tok_iter, iters, finish, wall, main=T) summary = { **_speed_summary(base, run, T.n_forward), "iterations": len(iters), "block": m, **_identity(base, run), } return {"runs": {"baseline": base, "method": run}, "summary": summary}