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()