#!/usr/bin/env python """load.py -- load a safetensors stage for benchmarking, correctly. python safetensors/load.py sft_7100 python safetensors/load.py base_62k --text "The capital of France is" Does the three things a harness has to get right (see README.md): * rebuilds the custom architecture from the repo's own model code * RE-TIES lm_head, which safetensors cannot store (shared storage) * reads max_position_embeddings from THAT stage's config, not a global """ import argparse import json import os import sys import torch from safetensors.torch import load_file HERE = os.path.dirname(os.path.abspath(__file__)) ROOT = os.path.dirname(HERE) sys.path.insert(0, ROOT) from config import SpikeWhaleConfig # noqa: E402 from model_v2 import SpikeWhaleLM # noqa: E402 from spike_tokenizer import SpikeTokenizer # noqa: E402 def load_stage(stage, device="cpu", dtype=torch.float32): d = stage if os.path.isdir(stage) else os.path.join(HERE, stage) cfg = json.load(open(os.path.join(d, "config.json"))) for k in ("architectures", "transformers_version", "dtype", "torch_dtype"): cfg.pop(k, None) model = SpikeWhaleLM(SpikeWhaleConfig(**cfg)) sd = load_file(os.path.join(d, "model.safetensors")) missing = model.load_state_dict(sd, strict=False).missing_keys if model.config.tie_word_embeddings: model.tie_weights() missing = [k for k in missing if k != "lm_head.weight"] if missing: raise RuntimeError(f"unexpected missing tensors: {missing[:8]}") model.eval().to(device=device, dtype=dtype) tok_path = os.path.join(d, "tokenizer.json") if not os.path.exists(tok_path): tok_path = os.path.join(ROOT, "tokenizer.json") return model, SpikeTokenizer(tok_path) def main(): ap = argparse.ArgumentParser() ap.add_argument("stage", help="base_62k | sft_7100 | dpo_3200 | a path") ap.add_argument("--device", default="cpu") ap.add_argument("--text", default="The capital of France is") args = ap.parse_args() model, tok = load_stage(args.stage, device=args.device) n = sum(p.numel() for p in {id(p): p for p in model.parameters()}.values()) print(f"[load] {args.stage}: {n / 1e6:.1f}M params | " f"ctx {model.config.max_position_embeddings} | " f"loop_count {getattr(model.config, 'loop_count', 1)} | " f"tied lm_head " f"{model.lm_head.weight.data_ptr() == model.model.embed_tokens.weight.data_ptr()}") ids = tok.encode(args.text) ids = ids.tolist() if hasattr(ids, "tolist") else list(ids) ids = [int(tok.bos_token_id)] + ids # training prepends with torch.no_grad(): out = model(torch.tensor([ids], device=args.device), use_cache=False) lg = out.logits[0, -1].float() top = lg.topk(5) print(f"[check] next-token top-5 after {args.text!r}:") for v, i in zip(top.values.tolist(), top.indices.tolist()): print(f" {tok.decode([i])!r:<14} {v:.3f}") labels = torch.tensor([ids], device=args.device) with torch.no_grad(): loss = model(labels, labels=labels, use_cache=False).loss print(f"[check] loss on the prompt itself: {float(loss):.4f}") if __name__ == "__main__": main()