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