"""Guided decoding (survey Def. 2): P'(w_t | w_ 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, }, }