Download tools/reference_sequence.py from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 3.87 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/reference_sequence.py
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/reference_sequence.py
-
curl -L -o reference_sequence.py https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/reference_sequence.py
3.87 kB
| #!/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() | |