Download decoding.py from usrnotfound101/anlp-a2-decoding: direct link, hf CLI and curl.
- Browser
- Download file 27.4 kB
-
https://huggingface.co/usrnotfound101/anlp-a2-decoding/resolve/main/decoding.py
- Command line
-
hf download hf://usrnotfound101/anlp-a2-decoding/decoding.py
-
curl -L -o decoding.py https://huggingface.co/usrnotfound101/anlp-a2-decoding/resolve/main/decoding.py
27.4 kB
| # ============================================================================= | |
| # decoding.py - Part 3: decoding strategies implemented from scratch on top of | |
| # EleutherAI/pythia-160m. Only model.forward() (i.e. model(...)) is used - | |
| # never model.generate(). | |
| # | |
| # greedy, beam search (width 1/2/4/...), top-k sampling, top-p (nucleus) sampling | |
| # | |
| # Evaluated on hamishivi/ROCStories (test split): prompt = first sentence, | |
| # reference = the remaining four sentences. | |
| # ============================================================================= | |
| import os, re, math, time, json, string | |
| from collections import Counter | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| DECODING_DEFAULTS = dict( | |
| model_name="EleutherAI/pythia-160m", | |
| dataset="hamishivi/ROCStories", split="test", n_samples=1000, seed=0, | |
| max_new_tokens=70, batch_size=32, beam_batch_size=8, dtype="float32", | |
| length_penalty=1.0, | |
| configs=[ | |
| {"name": "greedy", "strategy": "greedy"}, | |
| {"name": "beam_w1", "strategy": "beam", "width": 1}, | |
| {"name": "beam_w2", "strategy": "beam", "width": 2}, | |
| {"name": "beam_w4", "strategy": "beam", "width": 4}, | |
| {"name": "topk_k5", "strategy": "top_k", "k": 5}, | |
| {"name": "topk_k20", "strategy": "top_k", "k": 20}, | |
| {"name": "topk_k50", "strategy": "top_k", "k": 50}, | |
| {"name": "topk_k50_t0.7", "strategy": "top_k", "k": 50, "temperature": 0.7}, | |
| {"name": "topp_p0.5", "strategy": "top_p", "p": 0.5}, | |
| {"name": "topp_p0.8", "strategy": "top_p", "p": 0.8}, | |
| {"name": "topp_p0.95", "strategy": "top_p", "p": 0.95}, | |
| {"name": "topp_p0.9_t0.7", "strategy": "top_p", "p": 0.9, "temperature": 0.7}, | |
| ], | |
| timing_widths=[1, 2, 4, 8], timing_samples=100, | |
| judge_model="EleutherAI/pythia-410m", # external LM for an extra perplexity column (None to skip) | |
| use_bertscore=True, | |
| out_dir="/kaggle/working/outputs/decoding" if os.path.isdir("/kaggle/working") else "outputs/decoding", | |
| hf_user=None, repo_name="anlp-a2-decoding", push_to_hub=True, private=False, | |
| ) | |
| # ----------------------------------------------------------------------------- | |
| # model wrapper: one forward step with KV cache | |
| # ----------------------------------------------------------------------------- | |
| def load_lm(name, device, dtype="float32"): | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| tok = AutoTokenizer.from_pretrained(name) | |
| dt = {"float32": torch.float32, "float16": torch.float16, "bfloat16": torch.bfloat16}[dtype] | |
| try: | |
| model = AutoModelForCausalLM.from_pretrained(name, dtype=dt) | |
| except TypeError: | |
| model = AutoModelForCausalLM.from_pretrained(name, torch_dtype=dt) | |
| model.to(device).eval() | |
| if tok.pad_token_id is None: | |
| tok.pad_token = tok.eos_token | |
| return model, tok | |
| def reorder_cache(past, idx): | |
| """Select rows (beams) of a KV cache - works for legacy tuples and Cache objects.""" | |
| if past is None: | |
| return None | |
| if isinstance(past, (tuple, list)): | |
| return tuple(tuple(t.index_select(0, idx) for t in layer) for layer in past) | |
| if hasattr(past, "reorder_cache"): | |
| past.reorder_cache(idx) | |
| return past | |
| if hasattr(past, "layers"): # transformers >= 4.56 DynamicCache | |
| for layer in past.layers: | |
| layer.keys = layer.keys.index_select(0, idx) | |
| layer.values = layer.values.index_select(0, idx) | |
| return past | |
| if hasattr(past, "key_cache"): | |
| for i in range(len(past.key_cache)): | |
| past.key_cache[i] = past.key_cache[i].index_select(0, idx) | |
| past.value_cache[i] = past.value_cache[i].index_select(0, idx) | |
| return past | |
| raise TypeError(f"cannot reorder cache of type {type(past)}") | |
| def forward_step(model, input_ids, attention_mask, position_ids, past): | |
| out = model(input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, | |
| past_key_values=past, use_cache=True) | |
| return out.logits[:, -1, :].float(), out.past_key_values | |
| def left_pad(prompts, pad_id, device): | |
| L = max(len(p) for p in prompts) | |
| ids = torch.full((len(prompts), L), pad_id, dtype=torch.long) | |
| mask = torch.zeros((len(prompts), L), dtype=torch.long) | |
| for i, p in enumerate(prompts): | |
| ids[i, L - len(p):] = torch.tensor(p) | |
| mask[i, L - len(p):] = 1 | |
| pos = (mask.cumsum(1) - 1).clamp(min=0) | |
| return ids.to(device), mask.to(device), pos.to(device) | |
| # ----------------------------------------------------------------------------- | |
| # logit processors | |
| # ----------------------------------------------------------------------------- | |
| def top_k_filter(logits, k): | |
| k = min(k, logits.size(-1)) | |
| kth = torch.topk(logits, k, dim=-1).values[:, -1:] | |
| return logits.masked_fill(logits < kth, float("-inf")) | |
| def top_p_filter(logits, p): | |
| """Keep the smallest set of tokens whose cumulative probability >= p (always >= 1 token).""" | |
| sorted_logits, sorted_idx = torch.sort(logits, descending=True, dim=-1) | |
| probs = F.softmax(sorted_logits, dim=-1) | |
| cum = probs.cumsum(-1) | |
| remove = (cum - probs) >= p # mass *before* this token already reaches p | |
| remove[:, 0] = False | |
| sorted_logits = sorted_logits.masked_fill(remove, float("-inf")) | |
| return torch.full_like(logits, float("-inf")).scatter(-1, sorted_idx, sorted_logits) | |
| # ----------------------------------------------------------------------------- | |
| # greedy / top-k / top-p (batched) | |
| # ----------------------------------------------------------------------------- | |
| def sample_decode(model, prompts, max_new_tokens, eos_id, pad_id, strategy="greedy", k=None, p=None, | |
| temperature=1.0, generator=None): | |
| """Returns (list of generated id lists, list of sum log-prob under the *unmodified* | |
| model distribution, list of lengths).""" | |
| device = next(model.parameters()).device | |
| B = len(prompts) | |
| ids, mask, pos = left_pad(prompts, pad_id, device) | |
| logits, past = forward_step(model, ids, mask, pos, None) | |
| next_pos = pos[:, -1:] + 1 | |
| finished = torch.zeros(B, dtype=torch.bool, device=device) | |
| out, lp_sum, lengths = [], torch.zeros(B, device=device), torch.zeros(B, dtype=torch.long, device=device) | |
| for _ in range(max_new_tokens): | |
| logp = F.log_softmax(logits, -1) | |
| if strategy == "greedy": | |
| nxt = logits.argmax(-1) | |
| else: | |
| lg = logits / max(temperature, 1e-6) | |
| if strategy == "top_k": | |
| lg = top_k_filter(lg, k) | |
| elif strategy == "top_p": | |
| lg = top_p_filter(lg, p) | |
| nxt = torch.multinomial(F.softmax(lg, -1), 1, generator=generator).squeeze(1) | |
| nxt = torch.where(finished, torch.full_like(nxt, pad_id), nxt) | |
| lp_sum += torch.where(finished, torch.zeros_like(lp_sum), logp.gather(1, nxt[:, None]).squeeze(1)) | |
| lengths += (~finished).long() | |
| out.append(nxt) | |
| finished |= nxt == eos_id | |
| if bool(finished.all()): | |
| break | |
| mask = torch.cat([mask, torch.ones(B, 1, dtype=mask.dtype, device=device)], 1) | |
| logits, past = forward_step(model, nxt[:, None], mask, next_pos, past) | |
| next_pos = next_pos + 1 | |
| gen = torch.stack(out, 1).tolist() if out else [[] for _ in range(B)] | |
| res = [] | |
| for row in gen: | |
| if eos_id in row: | |
| row = row[:row.index(eos_id)] | |
| res.append(row) | |
| return res, lp_sum.tolist(), lengths.tolist() | |
| # ----------------------------------------------------------------------------- | |
| # beam search (batched over prompts, width W per prompt) | |
| # ----------------------------------------------------------------------------- | |
| def beam_search(model, prompts, width, max_new_tokens, eos_id, pad_id, length_penalty=1.0): | |
| """Standard beam search with a KV cache. | |
| * each step: scores(B,W) + log p(B,W,V) -> top 2W candidates per prompt | |
| * candidates ending in EOS become finished hypotheses (score / len^alpha) | |
| * a prompt is done once it has W finished hypotheses (early stopping) | |
| * at max length, open beams are added to the finished pool. | |
| Returns (ids, sum log-prob, length) of the best hypothesis per prompt.""" | |
| device = next(model.parameters()).device | |
| B, W = len(prompts), width | |
| rep = [p for p in prompts for _ in range(W)] | |
| ids, mask, pos = left_pad(rep, pad_id, device) | |
| logits, past = forward_step(model, ids, mask, pos, None) | |
| V = logits.size(-1) | |
| next_pos = pos[:, -1:] + 1 | |
| beam_scores = torch.zeros(B, W, device=device) | |
| beam_scores[:, 1:] = float("-inf") # all beams identical at t=0 -> only expand beam 0 | |
| seqs = torch.zeros(B * W, 0, dtype=torch.long, device=device) | |
| finished = [[] for _ in range(B)] # (normalised score, raw score, token list) | |
| done = [False] * B | |
| for t in range(max_new_tokens): | |
| logp = F.log_softmax(logits, -1).view(B, W, V) | |
| cand = (beam_scores.unsqueeze(-1) + logp).view(B, W * V) | |
| top_s, top_i = cand.topk(2 * W, dim=-1) | |
| top_s_l, top_i_l = top_s.tolist(), top_i.tolist() | |
| new_scores = torch.zeros(B, W, device=device) | |
| new_beam = torch.zeros(B, W, dtype=torch.long) | |
| new_tok = torch.full((B, W), pad_id, dtype=torch.long) | |
| seqs_l = None | |
| for b in range(B): | |
| if done[b]: | |
| new_scores[b] = float("-inf"); new_scores[b, 0] = 0.0 | |
| continue | |
| j = 0 | |
| for r in range(2 * W): | |
| s, ci = top_s_l[b][r], top_i_l[b][r] | |
| bi, ti = ci // V, ci % V | |
| if s == float("-inf"): | |
| break | |
| if ti == eos_id: | |
| if r < W: | |
| if seqs_l is None: | |
| seqs_l = seqs.tolist() | |
| toks = seqs_l[b * W + bi] | |
| n = len(toks) + 1 | |
| finished[b].append((s / (n ** length_penalty), s, toks, n)) | |
| continue | |
| new_scores[b, j] = s; new_beam[b, j] = bi; new_tok[b, j] = ti | |
| j += 1 | |
| if j == W: | |
| break | |
| while j < W: # not enough live candidates | |
| new_scores[b, j] = float("-inf"); j += 1 | |
| if len(finished[b]) >= W: | |
| done[b] = True | |
| if all(done): | |
| break | |
| gidx = (torch.arange(B).unsqueeze(1) * W + new_beam).view(-1).to(device) | |
| seqs = torch.cat([seqs.index_select(0, gidx), new_tok.view(-1, 1).to(device)], 1) | |
| beam_scores = new_scores | |
| past = reorder_cache(past, gidx) | |
| mask = torch.cat([mask.index_select(0, gidx), torch.ones(B * W, 1, dtype=mask.dtype, device=device)], 1) | |
| next_pos = next_pos.index_select(0, gidx) | |
| logits, past = forward_step(model, new_tok.view(-1, 1).to(device), mask, next_pos, past) | |
| next_pos = next_pos + 1 | |
| seqs_l = seqs.tolist(); sc = beam_scores.tolist() | |
| res, lps, lens = [], [], [] | |
| for b in range(B): | |
| if not done[b]: | |
| for w in range(W): | |
| if sc[b][w] > float("-inf"): | |
| toks = seqs_l[b * W + w] | |
| n = max(len(toks), 1) | |
| finished[b].append((sc[b][w] / (n ** length_penalty), sc[b][w], toks, n)) | |
| best = max(finished[b], key=lambda z: z[0]) | |
| res.append(best[2]); lps.append(best[1]); lens.append(best[3]) | |
| return res, lps, lens | |
| # ----------------------------------------------------------------------------- | |
| # metrics | |
| # ----------------------------------------------------------------------------- | |
| _PUNCT = re.compile(f"[{re.escape(string.punctuation)}]") | |
| def norm_words(s): | |
| return _PUNCT.sub(" ", s.lower()).split() | |
| def clean_generation(text, max_sentences=4): | |
| """Post-processing applied before scoring: cut at the first newline (Pythia tends to | |
| drift into unrelated text after a line break) and keep at most 4 sentences | |
| (ROCStories continuations are exactly 4 sentences).""" | |
| text = text.lstrip() | |
| text = text.split("\n")[0] | |
| ends = [m.end() for m in re.finditer(r"[.!?](?=\s|$)", text)] | |
| if len(ends) >= max_sentences: | |
| text = text[:ends[max_sentences - 1]] | |
| return " " + text.strip() | |
| def token_prf_acc(gen, ref): | |
| g, r = norm_words(gen), norm_words(ref) | |
| common = Counter(g) & Counter(r) | |
| ov = sum(common.values()) | |
| P = ov / len(g) if g else 0.0 | |
| R = ov / len(r) if r else 0.0 | |
| F1 = 2 * P * R / (P + R) if P + R > 0 else 0.0 | |
| acc = sum(1 for i in range(len(r)) if i < len(g) and g[i] == r[i]) / len(r) if r else 0.0 | |
| return P, R, F1, acc | |
| def distinct_n(texts, n): | |
| grams, total = set(), 0 | |
| for t in texts: | |
| w = norm_words(t) | |
| ng = [tuple(w[i:i + n]) for i in range(len(w) - n + 1)] | |
| grams.update(ng); total += len(ng) | |
| return len(grams) / max(total, 1) | |
| def repetition_rate(text, n=4): | |
| w = norm_words(text) | |
| ng = [tuple(w[i:i + n]) for i in range(len(w) - n + 1)] | |
| return 1 - len(set(ng)) / len(ng) if ng else 0.0 | |
| def continuation_ppl(model, prompt_ids, cont_ids, pad_id, batch_size=32): | |
| """Per-sample perplexity of cont given prompt (teacher forced).""" | |
| device = next(model.parameters()).device | |
| out = [] | |
| for s in range(0, len(prompt_ids), batch_size): | |
| P = prompt_ids[s:s + batch_size]; C = cont_ids[s:s + batch_size] | |
| seqs = [p + c for p, c in zip(P, C)] | |
| L = max(len(x) for x in seqs) | |
| ids = torch.full((len(seqs), L), pad_id, dtype=torch.long) | |
| att = torch.zeros((len(seqs), L), dtype=torch.long) | |
| lab = torch.full((len(seqs), L), -100, dtype=torch.long) | |
| for i, (p, c) in enumerate(zip(P, C)): | |
| x = p + c | |
| ids[i, :len(x)] = torch.tensor(x); att[i, :len(x)] = 1 | |
| lab[i, len(p):len(x)] = torch.tensor(c) if c else lab[i, len(p):len(x)] | |
| ids, att, lab = ids.to(device), att.to(device), lab.to(device) | |
| logits = model(input_ids=ids, attention_mask=att).logits.float() | |
| lp = F.cross_entropy(logits[:, :-1].reshape(-1, logits.size(-1)), lab[:, 1:].reshape(-1), | |
| ignore_index=-100, reduction="none").view(len(seqs), -1) | |
| n = (lab[:, 1:] != -100).sum(1).clamp(min=1) | |
| out += torch.exp(lp.sum(1) / n).tolist() | |
| return out | |
| def score_config(gens, refs, self_lp, self_len): | |
| import sacrebleu | |
| prf = np.array([token_prf_acc(g, r) for g, r in zip(gens, refs)]) | |
| ppl_self = [math.exp(-lp / max(n, 1)) for lp, n in zip(self_lp, self_len)] | |
| return { | |
| "precision": float(prf[:, 0].mean()), "recall": float(prf[:, 1].mean()), | |
| "f1": float(prf[:, 2].mean()), "accuracy": float(prf[:, 3].mean()), | |
| "bleu": sacrebleu.corpus_bleu(gens, [refs]).score, | |
| "self_ppl_mean": float(np.mean(ppl_self)), "self_ppl_median": float(np.median(ppl_self)), | |
| "distinct_1": distinct_n(gens, 1), "distinct_2": distinct_n(gens, 2), | |
| "rep_4gram": float(np.mean([repetition_rate(g) for g in gens])), | |
| "avg_words": float(np.mean([len(norm_words(g)) for g in gens])), | |
| } | |
| # ----------------------------------------------------------------------------- | |
| # driver | |
| # ----------------------------------------------------------------------------- | |
| def run_config(model, tok, prompts, conf, cfg, seed=0): | |
| device = next(model.parameters()).device | |
| eos, pad = tok.eos_token_id, tok.pad_token_id | |
| gen_ids, lps, lens = [None] * len(prompts), [None] * len(prompts), [None] * len(prompts) | |
| order = sorted(range(len(prompts)), key=lambda i: len(prompts[i])) | |
| g = torch.Generator(device=device); g.manual_seed(seed) | |
| bs = cfg["beam_batch_size"] if conf["strategy"] == "beam" else cfg["batch_size"] | |
| if device.type == "cuda": torch.cuda.synchronize() | |
| t0 = time.time() | |
| for s in range(0, len(order), bs): | |
| idx = order[s:s + bs] | |
| P = [prompts[i] for i in idx] | |
| if conf["strategy"] == "beam": | |
| r, lp, ln = beam_search(model, P, conf["width"], cfg["max_new_tokens"], eos, pad, cfg["length_penalty"]) | |
| else: | |
| r, lp, ln = sample_decode(model, P, cfg["max_new_tokens"], eos, pad, conf["strategy"], | |
| conf.get("k"), conf.get("p"), conf.get("temperature", 1.0), g) | |
| for i, a, b, c in zip(idx, r, lp, ln): | |
| gen_ids[i], lps[i], lens[i] = a, b, c | |
| if device.type == "cuda": torch.cuda.synchronize() | |
| return gen_ids, lps, lens, time.time() - t0 | |
| def beam_timing(model, tok, prompts, widths, max_new_tokens): | |
| """Wall time of beam search with batch size 1 (one prompt at a time) for each width. | |
| Two modes: 'natural' (stops early once W hypotheses hit EOS) and 'fixed' (EOS disabled, | |
| always max_new_tokens steps) - the fixed mode isolates the per-step cost as a function of W.""" | |
| device = next(model.parameters()).device | |
| res = [] | |
| beam_search(model, prompts[:2], 2, 5, tok.eos_token_id, tok.pad_token_id) # warm-up | |
| for mode, eos in (("natural", tok.eos_token_id), ("fixed", -1)): | |
| for w in widths: | |
| times, steps = [], [] | |
| for p in prompts: | |
| if device.type == "cuda": torch.cuda.synchronize() | |
| t0 = time.time() | |
| r, _, ln = beam_search(model, [p], w, max_new_tokens, eos, tok.pad_token_id) | |
| if device.type == "cuda": torch.cuda.synchronize() | |
| times.append(time.time() - t0); steps.append(max(len(r[0]), 1)) | |
| rec = {"mode": mode, "width": w, "mean_s_per_prompt": float(np.mean(times)), "std_s": float(np.std(times)), | |
| "mean_ms_per_token": float(np.mean([t / s for t, s in zip(times, steps)]) * 1e3), | |
| "mean_output_tokens": float(np.mean(steps)), "total_s": float(np.sum(times)), | |
| "n_prompts": len(prompts)} | |
| res.append(rec) | |
| print(f" [{mode}] width {w}: {rec['mean_s_per_prompt']*1e3:.1f} ms/prompt, " | |
| f"{rec['mean_ms_per_token']:.2f} ms/token, {rec['mean_output_tokens']:.1f} tokens") | |
| return res | |
| def run_decoding_eval(user_cfg=None): | |
| import sys | |
| sys.path.insert(0, ".") | |
| from common import get_hf_token, hf_login, hf_resolve_user, hf_push_folder, save_json | |
| cfg = dict(DECODING_DEFAULTS); cfg.update(user_cfg or {}) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| os.makedirs(cfg["out_dir"], exist_ok=True) | |
| torch.manual_seed(cfg["seed"]) | |
| token = get_hf_token(); hf_login(token) | |
| model, tok = load_lm(cfg["model_name"], device, cfg["dtype"]) | |
| print(f"loaded {cfg['model_name']} on {device}; eos={tok.eos_token_id} pad={tok.pad_token_id}") | |
| if cfg.get("dataset_override") is not None: | |
| rows = cfg["dataset_override"] | |
| else: | |
| from datasets import load_dataset | |
| rows = load_dataset(cfg["dataset"], split=cfg["split"]) | |
| prompts_txt = list(rows["prompt"]); refs_txt = list(rows["continuation"]) | |
| keep = [i for i, r in enumerate(refs_txt) if r and r.strip()] | |
| rng = np.random.default_rng(cfg["seed"]) | |
| keep = sorted(rng.choice(keep, min(cfg["n_samples"], len(keep)), replace=False).tolist()) | |
| prompts_txt = [prompts_txt[i] for i in keep]; refs_txt = [refs_txt[i] for i in keep] | |
| refs_txt = [" " + r.strip() for r in refs_txt] | |
| prompts = [tok(p)["input_ids"] for p in prompts_txt] | |
| print(f"{len(prompts)} test prompts") | |
| # reference perplexity (independent of the decoding strategy) | |
| ref_ids = [tok(r)["input_ids"] for r in refs_txt] | |
| ref_ppl = continuation_ppl(model, prompts, ref_ids, tok.pad_token_id) | |
| summary = {"reference_ppl_mean": float(np.mean(ref_ppl)), "reference_ppl_median": float(np.median(ref_ppl)), | |
| "n_samples": len(prompts), "model": cfg["model_name"], "max_new_tokens": cfg["max_new_tokens"]} | |
| print(f"reference continuation ppl: mean {summary['reference_ppl_mean']:.2f}, " | |
| f"median {summary['reference_ppl_median']:.2f}") | |
| results, generations = [], {"prompt": prompts_txt, "reference": refs_txt} | |
| for conf in cfg["configs"]: | |
| print(f"== {conf['name']} ==") | |
| ids, lps, lens, dt = run_config(model, tok, prompts, conf, cfg, cfg["seed"]) | |
| raw = [tok.decode(x, skip_special_tokens=True) for x in ids] | |
| gens = [clean_generation(t) for t in raw] | |
| m = score_config(gens, refs_txt, lps, lens) | |
| m.update({"name": conf["name"], "strategy": conf["strategy"], "params": conf, | |
| "time_s": dt, "ms_per_prompt": dt / len(prompts) * 1e3, | |
| "gen_tokens_mean": float(np.mean([len(x) for x in ids]))}) | |
| results.append(m) | |
| generations[conf["name"]] = gens | |
| generations[conf["name"] + "__raw"] = raw | |
| print(f" P {m['precision']:.3f} R {m['recall']:.3f} F1 {m['f1']:.3f} Acc {m['accuracy']:.3f} " | |
| f"BLEU {m['bleu']:.2f} self-ppl {m['self_ppl_median']:.2f} dist-2 {m['distinct_2']:.3f} " | |
| f"rep4 {m['rep_4gram']:.3f} time {dt:.1f}s") | |
| # sanity: beam width 1 must equal greedy | |
| if "greedy" in generations and "beam_w1" in generations: | |
| same = np.mean([a == b for a, b in zip(generations["greedy__raw"], generations["beam_w1__raw"])]) | |
| summary["beam_w1_equals_greedy_frac"] = float(same) | |
| print(f"beam(w=1) identical to greedy on {same*100:.1f}% of prompts") | |
| # external judge perplexity | |
| if cfg["judge_model"]: | |
| try: | |
| judge, _ = load_lm(cfg["judge_model"], device, cfg["dtype"]) | |
| for m in results: | |
| gen_ids = [tok(g)["input_ids"] for g in generations[m["name"]]] | |
| pp = continuation_ppl(judge, prompts, gen_ids, tok.pad_token_id) | |
| m["judge_ppl_median"] = float(np.median(pp)); m["judge_ppl_mean"] = float(np.mean(pp)) | |
| jr = continuation_ppl(judge, prompts, ref_ids, tok.pad_token_id) | |
| summary["reference_judge_ppl_median"] = float(np.median(jr)) | |
| del judge | |
| except Exception as e: | |
| print("[judge] skipped:", e) | |
| if cfg["use_bertscore"]: | |
| try: | |
| from bert_score import score as bscore | |
| for m in results: | |
| P, R, F1 = bscore(generations[m["name"]], refs_txt, lang="en", verbose=False, | |
| device=str(device), batch_size=64) | |
| m["bertscore_p"], m["bertscore_r"], m["bertscore_f1"] = (float(P.mean()), float(R.mean()), | |
| float(F1.mean())) | |
| except Exception as e: | |
| print("[bertscore] skipped:", e) | |
| print("beam-search timing (batch size 1):") | |
| timing = beam_timing(model, tok, prompts[:cfg["timing_samples"]], cfg["timing_widths"], cfg["max_new_tokens"]) | |
| save_json({"summary": summary, "results": results, "beam_timing": timing, "config": | |
| {k: v for k, v in cfg.items() if k != "dataset_override"}}, | |
| os.path.join(cfg["out_dir"], "decoding_results.json")) | |
| save_json(generations, os.path.join(cfg["out_dir"], "generations.json")) | |
| shutil_copy_code(cfg["out_dir"]) | |
| user = hf_resolve_user(cfg["hf_user"], token) | |
| if cfg["push_to_hub"] and user: | |
| hf_push_folder(cfg["out_dir"], f"{user}/{cfg['repo_name']}", token, cfg["private"], "decoding results") | |
| print(f"pushed to https://huggingface.co/{user}/{cfg['repo_name']}") | |
| return summary, results, timing, generations | |
| def shutil_copy_code(out_dir): | |
| import shutil | |
| for f in ("decoding.py", "common.py"): | |
| if os.path.exists(f): | |
| shutil.copy(f, os.path.join(out_dir, f)) | |
| # ----------------------------------------------------------------------------- | |
| # plots / tables for the write-up | |
| # ----------------------------------------------------------------------------- | |
| PALETTE = ["#2a78d6", "#eb6834", "#1baf7a", "#eda100", "#e87ba4", "#008300", "#4a3aa7", "#e34948"] | |
| STRATEGY_COLOR = {"greedy": PALETTE[0], "beam": PALETTE[6], "top_k": PALETTE[1], "top_p": PALETTE[2]} | |
| def decoding_report(res_json_path, out_dir): | |
| import pandas as pd | |
| import matplotlib.pyplot as plt | |
| from common import setup_plot_style | |
| setup_plot_style() | |
| R = json.load(open(res_json_path)) | |
| df = pd.DataFrame(R["results"]) | |
| cols = ["name", "accuracy", "precision", "recall", "f1", "bleu", "self_ppl_median", "judge_ppl_median", | |
| "bertscore_f1", "distinct_2", "rep_4gram", "avg_words", "ms_per_prompt"] | |
| table = df[[c for c in cols if c in df.columns]].copy() | |
| table.to_csv(os.path.join(out_dir, "decoding_table.csv"), index=False) | |
| colors = [STRATEGY_COLOR[s] for s in df["strategy"]] | |
| metrics = [("f1", "token F1"), ("accuracy", "positional accuracy"), ("bleu", "BLEU"), | |
| ("self_ppl_median", "self perplexity (median)"), ("distinct_2", "distinct-2"), | |
| ("rep_4gram", "repeated 4-gram rate")] | |
| fig, axes = plt.subplots(2, 3, figsize=(15, 7.5), constrained_layout=True) | |
| for ax, (k, lab) in zip(axes.flat, metrics): | |
| ax.bar(range(len(df)), df[k], color=colors, width=0.7) | |
| ax.set_xticks(range(len(df))); ax.set_xticklabels(df["name"], rotation=60, ha="right", fontsize=8) | |
| ax.set_title(lab); ax.grid(axis="x", visible=False) | |
| handles = [plt.Rectangle((0, 0), 1, 1, color=c) for c in STRATEGY_COLOR.values()] | |
| fig.legend(handles, STRATEGY_COLOR.keys(), loc="upper right", ncol=4) | |
| fig.savefig(os.path.join(out_dir, "decoding_metrics.png")) | |
| # quality vs diversity | |
| fig2, ax = plt.subplots(figsize=(6.5, 4.5), constrained_layout=True) | |
| for _, r in df.iterrows(): | |
| ax.scatter(r["distinct_2"], r["f1"], color=STRATEGY_COLOR[r["strategy"]], s=60, zorder=3, | |
| edgecolor="white", linewidth=1.5) | |
| ax.annotate(r["name"], (r["distinct_2"], r["f1"]), fontsize=7, xytext=(4, 3), textcoords="offset points") | |
| ax.set_xlabel("distinct-2 (diversity)"); ax.set_ylabel("token F1 vs reference") | |
| ax.set_title("quality / diversity trade-off") | |
| fig2.savefig(os.path.join(out_dir, "quality_vs_diversity.png")) | |
| # beam timing | |
| tm = pd.DataFrame(R["beam_timing"]) | |
| fig3, axes3 = plt.subplots(1, 2, figsize=(11, 3.8), constrained_layout=True) | |
| for i, (mode, g) in enumerate(tm.groupby("mode", sort=False)): | |
| axes3[0].plot(g["width"], g["mean_s_per_prompt"] * 1e3, marker="o", label=mode, color=PALETTE[i]) | |
| axes3[1].plot(g["width"], g["mean_ms_per_token"], marker="o", label=mode, color=PALETTE[i]) | |
| for ax, lab in zip(axes3, ["ms per prompt", "ms per generated token"]): | |
| ax.set_xlabel("beam width"); ax.set_ylabel(lab); ax.set_xticks(sorted(tm["width"].unique())); ax.legend() | |
| axes3[0].set_title("beam search wall time (batch size 1)") | |
| axes3[1].set_title("per-step cost") | |
| fig3.savefig(os.path.join(out_dir, "beam_timing.png")) | |
| return table, tm | |
| def show_examples(gen_json_path, n=5, names=None, seed=0): | |
| G = json.load(open(gen_json_path)) | |
| names = names or [k for k in G if k not in ("prompt", "reference") and not k.endswith("__raw")] | |
| rng = np.random.default_rng(seed) | |
| for i in rng.choice(len(G["prompt"]), n, replace=False): | |
| print("=" * 100) | |
| print("PROMPT :", G["prompt"][i]) | |
| print("REFERENCE:", G["reference"][i].strip()) | |
| for k in names: | |
| print(f"{k:>16s}: {G[k][i].strip()}") | |