Spaces:
Sleeping
Sleeping
Download decoding/contrastive.py from wang2226/beyond-tokens-decoding: direct link, hf CLI and curl.
- Browser
- Download file 10.9 kB
-
https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/decoding/contrastive.py
- Command line
-
hf download hf://spaces/wang2226/beyond-tokens-decoding/decoding/contrastive.py
-
curl -L -o contrastive.py https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/decoding/contrastive.py
10.9 kB
| """Contrastive decoding (survey Def. 1): P(w_t | w_<t) = softmax(z+ + alpha * (z+ - z-)). | |
| - Expert-amateur CD (Li et al. 2023): z+ from the main model, z- from the helper. | |
| - Two-prompt contrast with one model: CAD (Shi et al. 2024; with vs without context) | |
| and a ROSE-style preset (Zhong et al. 2024; normal vs reverse system prompt). | |
| - DoLa (Chuang et al. 2024): z+ from the final layer, z- from the premature layer | |
| whose early-exit distribution diverges most (JSD) from the final one. | |
| All scores are computed from log-probabilities, which differ from logits by a | |
| per-distribution constant, so softmax of the Def. 1 combination is unchanged. | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| from decoding.common import ( | |
| KV, | |
| ROUND, | |
| apc_mask, | |
| ar_decode, | |
| distinct_n, | |
| early_exit_logits, | |
| forward_stats, | |
| jsd, | |
| log_probs, | |
| make_generator, | |
| make_run, | |
| now, | |
| pick, | |
| rank_of, | |
| rep_penalty, | |
| topk_entries, | |
| ) | |
| from prompting import build_ids, prompt_of | |
| LENS_TOPK = 3 | |
| def contrast(z_pos: torch.Tensor, z_neg: torch.Tensor, alpha: float, beta: float) -> tuple[torch.Tensor, torch.Tensor]: | |
| """Def. 1 scores (1 + alpha) z+ - alpha z-, restricted to the plausible head of z+.""" | |
| head = apc_mask(z_pos, beta) | |
| s = (1 + alpha) * z_pos - alpha * z_neg | |
| return s.masked_fill(~head, float("-inf")), head | |
| def _greedy_or_sample(P: dict) -> tuple[bool, float]: | |
| return P.get("mode", "greedy") == "greedy", float(P.get("temperature", 1.0)) | |
| def _baseline(fam, model, prompt_ids, P, label, rep=None): | |
| greedy, temp = _greedy_or_sample(P) | |
| return ar_decode( | |
| model, | |
| fam.tok, | |
| prompt_ids, | |
| P["max_new_tokens"], | |
| fam.stop_ids, | |
| disp=fam.disp, | |
| greedy=greedy, | |
| temperature=temp, | |
| rep=P.get("rep", 1.0) if rep is None else rep, | |
| gen=make_generator(P.get("seed", 0)), | |
| label=label, | |
| ) | |
| def _pairwise_loop(fam, kv_pos, kv_neg, pos_prompt, neg_prompt, P, *, tau_neg=1.0): | |
| """Shared token loop for any Def. 1 contrast between two (model, prompt) pairs.""" | |
| greedy, temp = _greedy_or_sample(P) | |
| gen = make_generator(P.get("seed", 0)) | |
| alpha, beta, rep = float(P["alpha"]), float(P["beta"]), float(P.get("rep", 1.0)) | |
| gen_ids: list[int] = [] | |
| steps: list[dict] = [] | |
| finish = "length" | |
| for i in range(P["max_new_tokens"]): | |
| lp = log_probs(kv_pos.run(pos_prompt + gen_ids).logits[0, -1]) | |
| ln = log_probs(kv_neg.run(neg_prompt + gen_ids).logits[0, -1], tau_neg) | |
| s, head = contrast(lp, ln, alpha, beta) | |
| s = rep_penalty(s, pos_prompt + gen_ids, rep) | |
| t = pick(s, greedy, temp, gen) | |
| final = torch.softmax(s / (1.0 if greedy else max(temp, 1e-5)), dim=-1) | |
| top_pos = int(lp.argmax()) | |
| steps.append({ | |
| "i": i, | |
| "id": t, | |
| "changed": t != top_pos, | |
| "base_rank": rank_of(lp, t), | |
| "top": { | |
| "pos": topk_entries(lp.exp(), fam.disp), | |
| "neg": topk_entries(ln.exp(), fam.disp), | |
| "final": topk_entries(final, fam.disp), | |
| }, | |
| "x": { | |
| "head": int(head.sum()), | |
| "p_pos": round(float(lp[t].exp()), ROUND), | |
| "p_neg": round(float(ln[t].exp()), ROUND), | |
| "p_final": round(float(final[t]), ROUND), | |
| "pos_top": [top_pos, fam.disp(top_pos)], | |
| }, | |
| }) | |
| gen_ids.append(t) | |
| if t in fam.stop_ids: | |
| finish = "eos" | |
| break | |
| return gen_ids, steps, finish | |
| def _summary(base: dict, run: dict, third: dict | None) -> dict: | |
| steps = run["steps"] | |
| out = { | |
| "changed": sum(s["changed"] for s in steps), | |
| "changed_frac": round(sum(s["changed"] for s in steps) / max(len(steps), 1), 3), | |
| "distinct2": {"baseline": distinct_n(base["ids"]), "method": distinct_n(run["ids"])}, | |
| "length": {"baseline": len(base["ids"]), "method": len(run["ids"])}, | |
| } | |
| if third is not None: | |
| out["distinct2"]["third"] = distinct_n(third["ids"]) | |
| out["length"]["third"] = len(third["ids"]) | |
| return out | |
| def run_cd(fam, P: dict) -> dict: | |
| """Expert-amateur contrastive decoding: main model vs helper model.""" | |
| if fam.helper is None: | |
| raise ValueError(f"{fam.label} has no helper model, so expert-amateur CD is unavailable.") | |
| prompt = prompt_of(fam, P) | |
| main_name, helper_name = fam.main_id.split("/")[-1], fam.helper_id.split("/")[-1] | |
| base = _baseline(fam, fam.main, prompt, P, f"Expert alone ({main_name})") | |
| third = _baseline(fam, fam.helper, prompt, P, f"Amateur alone ({helper_name})") if P.get("show_third", True) else None | |
| E, A = KV(fam.main), KV(fam.helper) | |
| t0 = now(E.device) | |
| ids, steps, finish = _pairwise_loop(fam, E, A, prompt, prompt, P, tau_neg=float(P.get("tau", 1.0))) | |
| wall = now(E.device) - t0 | |
| run = make_run(f"Contrastive decoding (α={P['alpha']}, β={P['beta']})", fam.tok, ids, finish, steps=steps, | |
| **forward_stats(main=E, helper=A)) | |
| run["time_ms"]["wall"] = round(1000 * wall, 2) | |
| return { | |
| "prompt_tokens": {"main": len(prompt)}, | |
| "roles": {"pos": f"expert · {main_name}", "neg": f"amateur · {helper_name}"}, | |
| "runs": {"baseline": base, "method": run, "third": third}, | |
| "summary": _summary(base, run, third), | |
| } | |
| def run_two_prompt(fam, P: dict) -> dict: | |
| """One model, two prompts: CAD (with/without context) or ROSE-style (normal/reverse system prompt).""" | |
| raw = P.get("raw", False) | |
| if P["method"] == "cad": | |
| pos = build_ids(fam, P["prompt"], context=P["context"], raw=raw) | |
| neg = build_ids(fam, P["prompt"], raw=raw) | |
| roles = {"pos": "with context", "neg": "without context"} | |
| labels = ("Greedy with context", "Greedy without context", f"Context-aware decoding (α={P['alpha']})") | |
| else: | |
| pos = build_ids(fam, P["prompt"], system=P.get("system") or None, raw=raw) | |
| neg = build_ids(fam, P["prompt"], system=P["reverse_system"], raw=raw) | |
| roles = {"pos": "normal prompt", "neg": "reverse prompt"} | |
| labels = ("Greedy, normal prompt", "Greedy, reverse prompt", f"ROSE-style contrast (α={P['alpha']})") | |
| base = _baseline(fam, fam.main, pos, P, labels[0]) | |
| third = _baseline(fam, fam.main, neg, P, labels[1]) if P.get("show_third", True) else None | |
| K1, K2 = KV(fam.main), KV(fam.main) | |
| t0 = now(K1.device) | |
| ids, steps, finish = _pairwise_loop(fam, K1, K2, pos, neg, P) | |
| wall = now(K1.device) - t0 | |
| run = make_run(labels[2], fam.tok, ids, finish, steps=steps, **forward_stats(main=K1, main_neg=K2)) | |
| run["time_ms"]["wall"] = round(1000 * wall, 2) | |
| return { | |
| "prompt_tokens": {"pos": len(pos), "neg": len(neg)}, | |
| "roles": roles, | |
| "runs": {"baseline": base, "method": run, "third": third}, | |
| "summary": _summary(base, run, third), | |
| } | |
| def dola_candidates(n_layers: int, tied: bool, bucket: str) -> list[int]: | |
| """Candidate premature layers (indices into hidden_states), following HF's DoLa rule.""" | |
| if not tied: | |
| start = 0 | |
| elif n_layers > 2: | |
| start = 2 | |
| elif n_layers == 2: | |
| start = 1 | |
| else: | |
| start = 0 | |
| if bucket == "low": | |
| if start == n_layers // 2: | |
| return [start] | |
| return list(range(start, n_layers // 2, 2)) if n_layers <= 40 else list(range(start, 20, 2)) | |
| return list(range(n_layers // 2, n_layers, 2)) if n_layers <= 40 else list(range(n_layers - 20, n_layers, 2)) | |
| def run_dola(fam, P: dict) -> dict: | |
| """DoLa: contrast the final layer with the most divergent premature layer.""" | |
| prompt = prompt_of(fam, P) | |
| model = fam.main | |
| cfg = model.config.get_text_config() | |
| N = cfg.num_hidden_layers | |
| tied = bool(getattr(model.config, "tie_word_embeddings", False)) | |
| cands = dola_candidates(N, tied, P["bucket"]) | |
| greedy, temp = _greedy_or_sample(P) | |
| gen = make_generator(P.get("seed", 0)) | |
| beta, rep, norm_on = float(P["beta"]), float(P.get("rep", 1.2)), bool(P.get("apply_norm", True)) | |
| base = _baseline(fam, model, prompt, P, f"Greedy ({fam.main_id.split('/')[-1]})", rep=rep) | |
| kv = KV(model) | |
| seq = list(prompt) | |
| gen_ids: list[int] = [] | |
| steps: list[dict] = [] | |
| finish = "length" | |
| t0 = now(kv.device) | |
| for i in range(P["max_new_tokens"]): | |
| out = kv.run(seq, hidden=True) | |
| final = log_probs(out.logits[0, -1]) | |
| H = torch.stack([out.hidden_states[j][0, -1] for j in range(N)]) # j = 0 is the embedding output | |
| lens = log_probs(early_exit_logits(model, H, apply_norm=norm_on)) # [N, V] | |
| d = jsd(final, lens[cands]) | |
| M = cands[int(d.argmax())] | |
| s = (final - lens[M]).masked_fill(~apc_mask(final, beta), float("-inf")) | |
| s = rep_penalty(s, seq, rep) | |
| t = pick(s, greedy, temp, gen) | |
| contrast_p = torch.softmax(s / (1.0 if greedy else max(temp, 1e-5)), dim=-1) | |
| top_final = int(rep_penalty(final, seq, rep).argmax()) | |
| lens_p = lens.exp() | |
| steps.append({ | |
| "i": i, | |
| "id": t, | |
| "changed": t != top_final, | |
| "base_rank": rank_of(final, t), | |
| "top": { | |
| "pos": topk_entries(final.exp(), fam.disp), | |
| "neg": topk_entries(lens_p[M], fam.disp), | |
| "final": topk_entries(contrast_p, fam.disp), | |
| }, | |
| "x": { | |
| "premature": M, | |
| "jsd": [[j, round(float(v), 5)] for j, v in zip(cands, d.tolist())], | |
| "lens": [[j, topk_entries(lens_p[j], fam.disp, LENS_TOPK)] for j in range(1, N)] | |
| + [[N, topk_entries(final.exp(), fam.disp, LENS_TOPK)]], | |
| "p_pos": round(float(final[t].exp()), ROUND), | |
| "p_neg": round(float(lens_p[M][t]), ROUND), | |
| "p_final": round(float(contrast_p[t]), ROUND), | |
| "pos_top": [top_final, fam.disp(top_final)], | |
| }, | |
| }) | |
| seq.append(t) | |
| gen_ids.append(t) | |
| if t in fam.stop_ids: | |
| finish = "eos" | |
| break | |
| wall = now(kv.device) - t0 | |
| run = make_run(f"DoLa ({P['bucket']} layers, β={beta})", fam.tok, gen_ids, finish, steps=steps, | |
| **forward_stats(main=kv)) | |
| run["time_ms"]["wall"] = round(1000 * wall, 2) | |
| summary = _summary(base, run, None) | |
| premature = [s["x"]["premature"] for s in steps] | |
| summary["premature_counts"] = [[j, premature.count(j)] for j in cands] | |
| return { | |
| "prompt_tokens": {"main": len(prompt)}, | |
| "roles": {"pos": f"final layer {N}", "neg": "premature layer"}, | |
| "layers": {"n": N, "candidates": cands, "apply_norm": norm_on}, | |
| "runs": {"baseline": base, "method": run, "third": None}, | |
| "summary": summary, | |
| } | |