Ling-3.0-tiny-RKNN / tools /compare_mla_logits.py
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
2.53 kB
#!/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))