"""Load a released multilingual value-embedding GPT checkpoint into the portable GPT module in modeling_gpt.py. from load import load_model model = load_model("best_state.pt", size="medium") # eval mode, CPU logits = model.logits(input_ids) # (1, T, padded_vocab) nll = model.loss(input_ids, target_ids) """ from __future__ import annotations import json, os import torch from modeling_gpt import GPT, GPTConfig, CONFIGS _ARCH_KEYS = {"name", "variant", "num_layers", "num_heads", "model_dim", "head_dim", "vocab_size", "eot_token", "doc_mask_token", "long_window", "short_window", "max_seq_len"} def config_from_json(path: str) -> GPTConfig: d = json.load(open(path)) return GPTConfig(**{k: v for k, v in d.items() if k in _ARCH_KEYS}) def load_pretrained(model_dir: str, device: str = "cpu", max_seq_len: int = None) -> GPT: """Load a packaged model directory (best_state.pt + config.json).""" cfg = config_from_json(os.path.join(model_dir, "config.json")) ckpt = os.path.join(model_dir, "best_state.pt") return load_model(ckpt, cfg=cfg, device=device, max_seq_len=max_seq_len) def _strip(sd: dict) -> dict: # torch.compile wraps the module -> "_orig_mod." prefix (small/medium have it, # large does not). Remove it so keys match the plain module. return {k[len("_orig_mod."):] if k.startswith("_orig_mod.") else k: v for k, v in sd.items()} def load_model(ckpt_path: str, size: str = None, cfg: GPTConfig = None, device: str = "cpu", max_seq_len: int = None) -> GPT: ck = torch.load(ckpt_path, map_location="cpu", weights_only=False) sd = _strip(ck["model"]) if cfg is None: cfg = CONFIGS[size] if max_seq_len is not None: cfg.max_seq_len = max_seq_len model = GPT(cfg) # assign=True keeps each tensor's stored dtype (bf16 embeds / fp32 head) exactly. missing, unexpected = model.load_state_dict(sd, strict=False, assign=True) # rotary cos/sin are non-persistent buffers -> legitimately "missing"; nothing else should be. real_missing = [k for k in missing if ".rotary." not in k] assert not real_missing, f"missing params: {real_missing}" assert not unexpected, f"unexpected params: {unexpected}" return model.eval().to(device) if __name__ == "__main__": import sys m = load_model(sys.argv[1], size=sys.argv[2]) print(f"loaded {sys.argv[2]}: {sum(p.numel() for p in m.parameters())/1e6:.1f}M params; " f"module ready (eval).")