File size: 3,866 Bytes
3fd1a35 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 | #!/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()
|