ipa-gpt / load.py
mugezhang's picture
Slim checkpoints to model-only (drop step/val_loss/code); update README/loader
b9f2e1a verified
Raw
History Blame Contribute Delete
2.55 kB
"""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).")