wang2226's picture
Beyond Tokens decoding playground: contrastive, guided and parallel decoding
371d90c verified
Raw History Blame Contribute Delete
11.5 kB
"""Guided decoding (survey Def. 2): P'(w_t | w_<t) ∝ G(P(w_t | w_<t), C).
- Classifier-guided (RAD / FUDGE style): rescore the top-k candidates with an attribute
score of the text they would produce, optionally after a short greedy lookahead
(in the spirit of NeuroLogic A*esque).
- Grammar/format-constrained (DOMINO style): mask every token whose text would make the
output stop being a prefix of a string the regular expression accepts.
"""
from __future__ import annotations
import copy
import regex
import torch
from decoding.common import (
KV,
ROUND,
TOPK,
ar_decode,
log_probs,
make_generator,
make_run,
now,
pick,
rep_penalty,
sample,
topk_entries,
)
from prompting import prompt_of
REGEX_TIMEOUT = 0.05 # seconds per match call
MAX_PATTERN = 300
MAX_TIMEOUTS = 8
# ---------------------------------------------------------------------------
# Classifier-guided decoding
# ---------------------------------------------------------------------------
@torch.no_grad()
def rollout(kv: KV, cand: torch.Tensor, L: int) -> list[list[int]]:
"""Greedy continuations of length L for each candidate token, as one batch.
``kv.cache`` must cover exactly the current prefix. The cache is copied, so ``kv``
is left untouched.
"""
if L <= 0:
return [[] for _ in range(len(cand))]
cache = copy.deepcopy(kv.cache)
cache.batch_repeat_interleave(len(cand))
x = cand.view(-1, 1).to(kv.device)
outs = []
for _ in range(L):
logits = kv.model(input_ids=x, past_key_values=cache, use_cache=True, logits_to_keep=1).logits[:, -1]
x = logits.argmax(-1, keepdim=True)
outs.append(x)
return torch.cat(outs, dim=1).tolist()
def attribute_fn(P: dict, scorers: dict):
"""G's attribute score: log P(sentiment | text) or cosine(text, topic)."""
if P["attribute"] == "sentiment":
scorer, target = scorers["sentiment"], P["target"]
return lambda texts: scorer(texts, target)
scorer, topic = scorers["topic"], P["topic"]
return lambda texts: scorer(texts, topic)
def trajectory(tok, ids: list[int], score) -> list[float]:
"""Attribute score of every prefix of a generated text."""
if not ids:
return []
texts = [tok.decode(ids[: i + 1], skip_special_tokens=True) for i in range(len(ids))]
return [round(v, ROUND) for v in score(texts)]
@torch.no_grad()
def run_classifier(fam, P: dict, scorers: dict) -> dict:
prompt = prompt_of(fam, P)
greedy = P.get("mode", "greedy") == "greedy"
temp, lam, k, L = float(P.get("temperature", 1.0)), float(P["lam"]), int(P["k"]), int(P["lookahead"])
rep = float(P.get("rep", 1.0))
score = attribute_fn(P, scorers)
gen = make_generator(P.get("seed", 0))
base = ar_decode(fam.main, fam.tok, prompt, P["max_new_tokens"], fam.stop_ids, disp=fam.disp, greedy=greedy,
temperature=temp, rep=rep, gen=make_generator(P.get("seed", 0)),
label=f"Unguided ({fam.main_id.split('/')[-1]})")
kv = KV(fam.main)
gen_ids: list[int] = []
steps: list[dict] = []
rollout_forwards = 0
finish = "length"
t0 = now(kv.device)
for i in range(P["max_new_tokens"]):
z = rep_penalty(kv.run(prompt + gen_ids).logits[0, -1].float(), prompt + gen_ids, rep)
lp = log_probs(z)
top_lp, cand = lp.topk(min(k, lp.shape[-1]))
conts = rollout(kv, cand, L)
rollout_forwards += L
cand_ids = cand.tolist()
texts = [fam.tok.decode(gen_ids + [c] + r, skip_special_tokens=True) for c, r in zip(cand_ids, conts)]
attr = torch.tensor(score(texts), dtype=torch.float32)
combined = top_lp.float().cpu() + lam * attr
j = int(combined.argmax()) if greedy else sample(torch.softmax(combined / max(temp, 1e-5), -1), gen)
t = cand_ids[j]
steps.append({
"i": i,
"id": t,
"changed": j != 0,
"base_rank": j,
"top": {"base": topk_entries(lp.exp(), fam.disp)},
"x": {
"cands": [
[c, fam.disp(c), round(float(l), ROUND), round(float(a), ROUND), round(float(s), ROUND),
fam.tok.decode(r, skip_special_tokens=True)]
for c, l, a, s, r in zip(cand_ids, top_lp.tolist(), attr.tolist(), combined.tolist(), conts)
],
"chosen": j,
},
})
gen_ids.append(t)
if t in fam.stop_ids:
finish = "eos"
break
wall = now(kv.device) - t0
run = make_run(f"Classifier-guided (λ={lam}, k={k}, L={L})", fam.tok, gen_ids, finish, steps=steps,
n_forward={"main": kv.n_forward, "rollout": rollout_forwards},
time_ms={"wall": round(1000 * wall, 2)})
traj = {"baseline": trajectory(fam.tok, base["ids"], score), "method": trajectory(fam.tok, gen_ids, score)}
final = {key: (v[-1] if v else None) for key, v in traj.items()}
return {
"prompt_tokens": {"main": len(prompt)},
"runs": {"baseline": base, "method": run, "third": None},
"trajectory": traj,
"summary": {
"attribute": P["attribute"],
"target": P["target"] if P["attribute"] == "sentiment" else P["topic"],
"score_kind": "log P(target | text)" if P["attribute"] == "sentiment" else "cosine(text, topic)",
"attr_final": final,
"changed": sum(s["changed"] for s in steps),
"changed_frac": round(sum(s["changed"] for s in steps) / max(len(steps), 1), 3),
},
}
# ---------------------------------------------------------------------------
# Grammar / format-constrained decoding
# ---------------------------------------------------------------------------
class Constraint:
"""Incremental regex check over decoded text, with timeouts against pathological patterns."""
def __init__(self, pattern: str):
if len(pattern) > MAX_PATTERN:
raise ValueError(f"Pattern is longer than {MAX_PATTERN} characters.")
try:
self.pat = regex.compile(pattern)
except regex.error as exc:
raise ValueError(f"Invalid regular expression: {exc}") from exc
self.timeouts = 0
def _match(self, text: str, partial: bool) -> bool:
try:
return self.pat.fullmatch(text, partial=partial, timeout=REGEX_TIMEOUT) is not None
except TimeoutError:
self.timeouts += 1
if self.timeouts > MAX_TIMEOUTS:
raise ValueError("This pattern is too slow to match incrementally; please simplify it.")
return False
def complete(self, text: str) -> bool:
return self._match(text, partial=False)
def viable(self, text: str) -> bool:
return self._match(text, partial=True)
def contains(self, text: str) -> bool:
try:
return self.pat.search(text, timeout=REGEX_TIMEOUT) is not None
except TimeoutError:
return False
def _token_ok(fam, con: Constraint, gen_ids: list[int], text: str, c: int, complete: bool, special: set[int]) -> bool:
if c in fam.stop_ids:
return complete # end-of-sequence only once the whole pattern has matched
if c in special:
return False
s = fam.tok.decode(gen_ids + [c], skip_special_tokens=False)
if s == text or "�" in s[len(text):]:
return False
return con.viable(s)
def _scan_vocab(fam, con, gen_ids, text, lp, complete, special, skip: set[int], want: int) -> list[int]:
"""Fallback: walk the whole vocabulary in probability order until ``want`` tokens pass."""
found: list[int] = []
for c in torch.argsort(lp, descending=True).tolist():
if c in skip:
continue
piece = fam.vocab_strings[c] if c < len(fam.vocab_strings) else ""
if c not in fam.stop_ids and (not piece or not con.viable(text + piece)):
continue # cheap approximate filter before the exact decode check
if _token_ok(fam, con, gen_ids, text, c, complete, special):
found.append(c)
if len(found) >= want:
break
return found
@torch.no_grad()
def run_regex(fam, P: dict) -> dict:
prompt = prompt_of(fam, P)
con = Constraint(P["pattern"])
greedy = P.get("mode", "greedy") == "greedy"
temp, K = float(P.get("temperature", 1.0)), int(P["top_k"])
gen = make_generator(P.get("seed", 0))
special = set(fam.tok.all_special_ids) - set(fam.stop_ids)
base = ar_decode(fam.main, fam.tok, prompt, P["max_new_tokens"], fam.stop_ids, disp=fam.disp, greedy=greedy,
temperature=temp, gen=make_generator(P.get("seed", 0)),
label=f"Unconstrained ({fam.main_id.split('/')[-1]})")
kv = KV(fam.main)
gen_ids: list[int] = []
text = ""
steps: list[dict] = []
finish = "length"
t0 = now(kv.device)
for i in range(P["max_new_tokens"]):
lp = log_probs(kv.run(prompt + gen_ids).logits[0, -1])
probs = lp.exp()
complete = con.complete(text)
top = lp.topk(min(K, lp.shape[-1])).indices.tolist()
ok = {c: _token_ok(fam, con, gen_ids, text, c, complete, special) for c in top}
allowed = [c for c in top if ok[c]]
fallback = False
if not allowed:
fallback = True
allowed = _scan_vocab(fam, con, gen_ids, text, lp, complete, special, set(top), 1 if greedy else 16)
if not allowed:
finish = "complete" if complete else "dead_end"
break
if greedy:
t = allowed[0] # candidates are in probability order
else:
idx = torch.tensor(allowed)
t = allowed[sample(torch.softmax(lp[idx.to(lp.device)] / max(temp, 1e-5), -1), gen)]
top1 = top[0]
steps.append({
"i": i,
"id": t,
"changed": t != top1,
"base_rank": int((lp > lp[t]).sum()),
"top": {"base": topk_entries(probs, fam.disp)},
"x": {
"checks": [[c, fam.disp(c), round(float(probs[c]), ROUND), ok[c]] for c in top[:TOPK]],
"allowed": sum(ok.values()),
"allowed_mass": round(float(sum(probs[c] for c in top if ok[c])), ROUND),
"fallback": fallback,
"complete": complete,
"p": round(float(probs[t]), ROUND),
},
})
gen_ids.append(t)
if t in fam.stop_ids:
finish = "eos"
break
text = fam.tok.decode(gen_ids, skip_special_tokens=False)
wall = now(kv.device) - t0
run = make_run("Regex-constrained", fam.tok, gen_ids, finish, steps=steps, n_forward={"main": kv.n_forward},
time_ms={"wall": round(1000 * wall, 2)})
base_text = base["text"].strip()
method_text = run["text"]
return {
"prompt_tokens": {"main": len(prompt)},
"runs": {"baseline": base, "method": run, "third": None},
"summary": {
"pattern": P["pattern"],
"valid": {"baseline": con.complete(base_text), "method": con.complete(method_text)},
"contains": {"baseline": con.contains(base_text)},
"changed": sum(s["changed"] for s in steps),
"fallback_steps": sum(s["x"]["fallback"] for s in steps),
"finish": finish,
},
}