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