"""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 # noqa: E402 from tokenizers import Tokenizer # noqa: E402 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 # 4 fixed tokens surround the state 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|>']]