| """Shared loading / inference helpers for the released Lean 4 tactic model. |
| |
| Everything is relative to the package root, so the package works as-is wherever |
| it is unpacked. A "variant" is a directory under checkpoints/ holding |
| lit_model.pth (and usually model.safetensors). |
| """ |
| import json |
| import os |
| import sys |
|
|
| import torch |
|
|
| HERE = os.path.dirname(os.path.abspath(__file__)) |
| ROOT = os.path.abspath(os.path.join(HERE, os.pardir)) |
| sys.path.insert(0, HERE) |
|
|
| from leanoar.model import ModelConfig, SchemeC |
| from tokenizers import Tokenizer |
|
|
| CFG_PATH = os.path.join(ROOT, 'config.json') |
| VARIANTS = ('checkpoints/e6', 'checkpoints/stage3-e3') |
|
|
|
|
| def load_config(): |
| return json.load(open(CFG_PATH)) |
|
|
|
|
| def specials(): |
| s = json.load(open(os.path.join(ROOT, 'special_tokens_v1.json'))) |
| return {str(k): int(v) for k, v in s.items()} if isinstance(s, dict) \ |
| else {str(i): int(v) for i, v in enumerate(s)} |
|
|
|
|
| def whitelist(): |
| return [w['id'] for w in json.load(open(os.path.join(ROOT, 'first_token_whitelist.json')))] |
|
|
|
|
| def pick_device(device=None): |
| if device: |
| return device |
| return 'cuda' if torch.cuda.is_available() else 'cpu' |
|
|
|
|
| def resolve_ckpt(ckpt='checkpoints/e6'): |
| """Accept a variant dir ('checkpoints/e6') or a direct path to a .pth/.safetensors.""" |
| if not os.path.isabs(ckpt) and not ckpt.startswith('checkpoints'): |
| ckpt = os.path.join(ROOT, ckpt) |
| if os.path.isdir(ckpt): |
| return ckpt, os.path.join(ckpt, 'lit_model.pth'), os.path.join(ckpt, 'model.safetensors') |
| return os.path.dirname(ckpt), ckpt, ckpt.replace('.pth', '.safetensors') |
|
|
|
|
| def load_model(ckpt='checkpoints/e6', device=None): |
| """Build SchemeC from config.json and load a packaged variant. Returns (net, cfg, device).""" |
| cfg_d = load_config() |
| cfg = ModelConfig(**{k: v for k, v in cfg_d['arch'].items() |
| if k in ModelConfig.__dataclass_fields__}) |
| device = pick_device(device) |
| dirname, pth, st = resolve_ckpt(ckpt) |
| if not (os.path.exists(pth) or os.path.exists(st)): |
| raise SystemExit(f'no weights found for {ckpt!r} (looked for {pth} / {st})') |
| net = SchemeC(cfg).to(device).eval() |
| if os.path.exists(st): |
| from safetensors.torch import load_file |
| net.load_state_dict(load_file(st)) |
| else: |
| net.load_state_dict(torch.load(pth, map_location=device, weights_only=False)['model']) |
| return net, cfg_d, device |
|
|
|
|
| def load_tokenizer(): |
| return Tokenizer.from_file(os.path.join(ROOT, 'tokenizer_v1.json')) |
|
|
|
|
| def encode_state(tok, state, cfg=None, max_new_tokens=None): |
| """[bos, <|FORWARD|>, <|state|>] + state tokens + [<|tactic|>]; keeps head+tail if too long.""" |
| cfg = cfg or load_config() |
| sp = specials() |
| ctx = cfg['context_length'] |
| mnt = max_new_tokens or cfg['prompt_template']['max_new_tokens'] |
| sids = tok.encode(state).ids |
| budget = ctx - 4 - mnt |
| if len(sids) > budget: |
| head = sids[:int(budget * 0.6)] |
| sids = head + sids[-(budget - len(head)):] |
| return [1, sp['<|FORWARD|>'], sp['<|state|>']] + list(sids) + [sp['<|tactic|>']] |
|
|