Download tools/compare_mla_logits.py from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 2.53 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/compare_mla_logits.py
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/compare_mla_logits.py
-
curl -L -o compare_mla_logits.py https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/compare_mla_logits.py
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)) | |