| """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: |
| |
| |
| 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) |
| |
| missing, unexpected = model.load_state_dict(sd, strict=False, assign=True) |
| |
| 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).") |
|
|