#!/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))