#!/usr/bin/env python3 """Compare teacher-forced distributions; raw cosine also includes irrelevant logit offsets.""" import argparse import array import json import math from pathlib import Path p = argparse.ArgumentParser(description=__doc__) p.add_argument("reference", type=Path) p.add_argument("candidate", type=Path) p.add_argument("--vocab", type=int, default=157184) p.add_argument("--teacher", type=Path) a = p.parse_args() size = a.reference.stat().st_size if size != a.candidate.stat().st_size or size % (a.vocab * 4): raise SystemExit("logit file shape mismatch") rows = [] teacher = list(map(int,a.teacher.read_text().split())) if a.teacher else None with a.reference.open("rb") as ref, a.candidate.open("rb") as new: for step in range(size // (a.vocab * 4)): x, y = array.array("f"), array.array("f") x.fromfile(ref, a.vocab) y.fromfile(new, a.vocab) if not all(math.isfinite(v) for v in x) or not all(math.isfinite(v) for v in y): raise SystemExit("nonfinite logits") sx = max(x) sy = max(y) sx += math.log(sum(math.exp(v-sx) for v in x)) sy += math.log(sum(math.exp(v-sy) for v in y)) kl = tv = 0.0 for u, v in zip(x, y): px, py = math.exp(u-sx), math.exp(v-sy) kl += px * ((u-sx)-(v-sy)) tv += abs(px-py) / 2 tx = max(range(a.vocab), key=x.__getitem__) ty = max(range(a.vocab), key=y.__getitem__) rows.append({"step": step, "kl_ref_candidate": kl, "total_variation": tv, "top1_equal": tx == ty, "ref_token": tx, "candidate_token": ty, "ref_margin_to_candidate": x[tx]-x[ty], "candidate_margin_to_ref": y[ty]-y[tx], "bitwise_equal": x == y}) if teacher: rows[-1]["reference_nll"] = sx-x[teacher[step]] rows[-1]["candidate_nll"] = sy-y[teacher[step]] extra = {} if teacher: rn=sum(r["reference_nll"] for r in rows)/len(rows) cn=sum(r["candidate_nll"] for r in rows)/len(rows) extra={"reference_mean_nll":rn,"candidate_mean_nll":cn,"perplexity_ratio":math.exp(cn-rn)} print(json.dumps({"reference": str(a.reference), "candidate": str(a.candidate), "mean_kl": sum(r["kl_ref_candidate"] for r in rows)/len(rows), "mean_total_variation": sum(r["total_variation"] for r in rows)/len(rows), "top1_agreement": sum(r["top1_equal"] for r in rows), **extra, "rows": rows}, indent=2))