anlp-a2-decoding / decoding.py
usrnotfound101's picture
decoding results
48a7823 verified
Raw History Blame Contribute Delete
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)}")
@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()}")