Spaces:
Running on Zero
Running on Zero
Download decoding/parallel.py from wang2226/beyond-tokens-decoding: direct link, hf CLI and curl.
- Browser
- Download file 14.9 kB
-
https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/decoding/parallel.py
- Command line
-
hf download hf://spaces/wang2226/beyond-tokens-decoding/decoding/parallel.py
-
curl -L -o parallel.py https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/decoding/parallel.py
14.9 kB
| """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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| 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 | |
| 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} | |