# ============================================================================= # 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)}") @torch.no_grad() 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) # ----------------------------------------------------------------------------- @torch.no_grad() 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) # ----------------------------------------------------------------------------- @torch.no_grad() 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 @torch.no_grad() 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()}")