Download tools/score_logits.py from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 1.57 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/score_logits.py
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/score_logits.py
-
curl -L -o score_logits.py https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/score_logits.py
1.57 kB
| #!/usr/bin/env python3 | |
| """Score saved logits against the same official teacher-forced reference.""" | |
| import argparse | |
| import json | |
| import numpy as np | |
| import torch | |
| p = argparse.ArgumentParser(description=__doc__) | |
| p.add_argument("reference") | |
| p.add_argument("--teacher-ids", required=True) | |
| p.add_argument("--candidate", action="append", required=True, help="NAME:FILE.f32") | |
| p.add_argument("--vocab", type=int, default=157184) | |
| args = p.parse_args() | |
| torch.set_num_threads(4) | |
| teacher = [int(x) for x in open(args.teacher_ids).read().split()] | |
| def read(path): | |
| result = torch.from_numpy(np.fromfile(path, dtype=np.float32).reshape(len(teacher), args.vocab)).double() | |
| if not torch.isfinite(result).all(): | |
| raise ValueError(path) | |
| return result | |
| reference = read(args.reference) | |
| log_p = reference.log_softmax(-1) | |
| indices = torch.arange(len(teacher)) | |
| summary = {"steps": len(teacher), "reference_nll": -log_p[indices, teacher].mean().item(), "candidates": {}} | |
| for item in args.candidate: | |
| name, path = item.split(":", 1) | |
| actual = read(path) | |
| log_q = actual.log_softmax(-1) | |
| cosine = torch.nn.functional.cosine_similarity(reference, actual, dim=-1) | |
| kl = (log_p.exp() * (log_p - log_q)).sum(-1) | |
| summary["candidates"][name] = { | |
| "nll": -log_q[indices, teacher].mean().item(), "kl_mean": kl.mean().item(), | |
| "kl_max": kl.max().item(), "cosine_mean": cosine.mean().item(), | |
| "cosine_min": cosine.min().item(), | |
| "top1_agreement": int((actual.argmax(-1) == reference.argmax(-1)).sum()), | |
| } | |
| print(json.dumps(summary, indent=2)) | |