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