wang2226's picture
Beyond Tokens decoding playground: contrastive, guided and parallel decoding
371d90c verified
Raw History Blame Contribute Delete
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
# ---------------------------------------------------------------------------
@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}