beyond-tokens-decoding / decoding /contrastive.py
wang2226's picture
Beyond Tokens decoding playground: contrastive, guided and parallel decoding
371d90c verified
Raw History Blame Contribute Delete
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
@torch.no_grad()
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),
}
@torch.no_grad()
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))
@torch.no_grad()
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,
}