#!/usr/bin/env python3 """Official checkpoint/implementation teacher-forced logits, on a CUDA host. No model downloads. Input IDs come from the C++ engine. This is a numerical reference, not a held-out language-quality benchmark: continuation IDs may have been generated by the old quantized engine. """ import argparse import json from pathlib import Path import numpy as np import torch from safetensors import safe_open from transformers import AutoModelForCausalLM def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--source", type=Path, required=True) parser.add_argument("--prompt-ids", type=Path, required=True) parser.add_argument("--continuation-ids", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--candidate", action="append", default=[], help="NAME:LOGITS.f32") args = parser.parse_args() torch.set_num_threads(4) torch.backends.cuda.matmul.allow_tf32 = False prompt = [int(x) for x in args.prompt_ids.read_text().split()] continuation = [int(x) for x in args.continuation_ids.read_text().split()] assert prompt and continuation model = AutoModelForCausalLM.from_pretrained( args.source, trust_remote_code=True, local_files_only=True, torch_dtype=torch.bfloat16, device_map={"": "cuda"}, attn_implementation="eager", ).eval() # Loading with torch_dtype may cast source FP32 gate parameters. Preserve # the original checkpoint's FP32 tensors instead of silently weakening it. weight_map = json.loads((args.source / "model.safetensors.index.json").read_text())["weight_map"] parameters = dict(model.named_parameters()) for shard in sorted(set(weight_map.values())): with safe_open(args.source / shard, framework="pt", device="cpu") as source: for name in source.keys(): if name in parameters and source.get_slice(name).get_dtype() == "F32": parameters[name].data = source.get_tensor(name).to("cuda") ids = torch.tensor([prompt + continuation[:-1]], device="cuda") with torch.inference_mode(): result = model(input_ids=ids, use_cache=False) logits = result.logits[0, len(prompt) - 1:].float().cpu() args.output.parent.mkdir(parents=True, exist_ok=True) logits.numpy().tofile(args.output) reference = logits.double() log_p = torch.log_softmax(reference, dim=-1) p = log_p.exp() summary = {"prompt_tokens": len(prompt), "steps": len(continuation), "vocab": logits.shape[-1], "reference": "official checkpoint and model code, BF16, source FP32 parameters restored", "candidates": {}} for item in args.candidate: name, path = item.split(":", 1) candidate = torch.from_numpy(np.fromfile(path, dtype=np.float32).reshape(logits.shape)).double() if not torch.isfinite(candidate).all(): raise ValueError(f"non-finite candidate {name}") cosine = torch.nn.functional.cosine_similarity(reference, candidate, dim=-1) kl = (p * (log_p - torch.log_softmax(candidate, dim=-1))).sum(-1) summary["candidates"][name] = { "cosine_mean": cosine.mean().item(), "cosine_min": cosine.min().item(), "kl_mean": kl.mean().item(), "kl_max": kl.max().item(), "top1_agreement": (reference.argmax(-1) == candidate.argmax(-1)).sum().item(), "teacher_nll": -torch.log_softmax(candidate, dim=-1)[torch.arange(len(continuation)), continuation].mean().item(), } summary["reference_teacher_nll"] = -log_p[torch.arange(len(continuation)), continuation].mean().item() args.output.with_suffix(".json").write_text(json.dumps(summary, indent=2) + "\n") print(json.dumps(summary, indent=2), flush=True) if __name__ == "__main__": main()