| |
| """Minimal inference for the released Lean 4 tactic model: proof state -> top-k tactics. |
| |
| python generate.py # real dev example #0 (packaged) |
| python generate.py --example 3 --k 8 |
| python generate.py --state-file s.txt # state from a Lean `unsolved goals` message |
| python generate.py --ckpt checkpoints/stage3-e3 # use the other packaged variant |
| |
| Needs torch, tokenizers, numpy (+ safetensors if the export exists). |
| The model proposes single tactics; only the Lean kernel can say whether they are correct. |
| """ |
| import argparse |
| import json |
| import os |
|
|
| import torch |
| import torch.nn.functional as F |
|
|
| from common import ROOT, load_config, load_model, load_tokenizer, encode_state, specials, \ |
| whitelist |
|
|
| CFG = load_config() |
| MAX_NEW = CFG['prompt_template']['max_new_tokens'] |
| SP = specials() |
| DANGLING = ('by', 'at', 'with', 'using', 'from', 'in', ',', ';', ':', '=>') |
|
|
|
|
| @torch.no_grad() |
| def propose(net, tok, wl, state, k=5, max_new=MAX_NEW, device='cpu'): |
| """Top-k single-tactic proposals: top-k first tokens, then greedy continuation. |
| |
| The first token is restricted to the shipped whitelist (97.6% of the first tokens in the |
| training split) - the one FSM constraint that makes candidates worth verifying. |
| Returns [(tactic, avg_logprob_of_its_tokens)]. |
| """ |
| p = encode_state(tok, state, CFG, max_new_tokens=max_new) |
| ids = torch.tensor([p], device=device) |
| logits = net(ids)['logits'][0, -1] |
| allow = torch.full_like(logits, float('-inf')) |
| allow[torch.tensor(wl, device=device)] = 0.0 |
| lp_all = F.log_softmax((logits + allow).float(), -1) |
| top = torch.topk(logits + allow, min(k, len(wl))).indices |
| firsts, scores = top.tolist(), lp_all[top].tolist() |
| seqs = torch.cat([ids.repeat(len(firsts), 1), |
| torch.tensor(firsts, device=device)[:, None]], dim=1) |
| done = [False] * len(firsts) |
| for _ in range(max_new): |
| if seqs.shape[1] > CFG['context_length']: |
| break |
| lg = net(seqs)['logits'][:, -1] |
| nxt = lg.argmax(-1) |
| step_lp = F.log_softmax(lg.float(), -1)[torch.arange(len(firsts), device=device), nxt] |
| for j in range(len(firsts)): |
| if not done[j]: |
| scores[j] += float(step_lp[j]) |
| done[j] = int(nxt[j]) == SP['<|eos|>'] |
| seqs = torch.cat([seqs, nxt[:, None]], dim=1) |
| if all(done): |
| break |
| out = [] |
| for j, gen in enumerate(seqs[:, len(p):].tolist()): |
| if gen and gen[-1] == SP['<|eos|>']: |
| gen = gen[:-1] |
| |
| |
| txt = tok.decode(gen).strip().split(' ')[0].strip() |
| while txt.split() and txt.split()[-1] in DANGLING: |
| txt = ' '.join(txt.split()[:-1]) |
| if txt and txt not in [c[0] for c in out]: |
| out.append((txt, scores[j] / max(1, len(gen)))) |
| return out |
|
|
|
|
| if __name__ == '__main__': |
| ap = argparse.ArgumentParser() |
| ap.add_argument('--ckpt', default='checkpoints/e6') |
| ap.add_argument('--state-file', help='file holding a proof state') |
| ap.add_argument('--example', type=int, |
| help='index into examples/dev_sample.jsonl (16 real dev states)') |
| ap.add_argument('--k', type=int, default=5) |
| ap.add_argument('--device', default=None) |
| a = ap.parse_args() |
|
|
| net, _cfg, device = load_model(a.ckpt, a.device) |
| tok = load_tokenizer() |
| wl = whitelist() |
| truth = None |
| if a.state_file: |
| state = open(a.state_file).read() |
| else: |
| idx = 0 if a.example is None else a.example |
| rec = [json.loads(l) for l in |
| open(os.path.join(ROOT, 'examples/dev_sample.jsonl'))][idx] |
| state, truth = rec['state'], rec['true_tactic'] |
| print(f'--- real dev example #{idx} (row {rec["dev_row"]} of the packaged dev split) ---') |
| print('--- state ---') |
| print(state.strip()) |
| print(f'--- top-{a.k} tactics (device={device}, ckpt={a.ckpt}) ---') |
| for t, lp in propose(net, tok, wl, state, a.k, device=device): |
| print(f' {lp:+.3f} {t}') |
| if truth: |
| print('--- the tactic mathlib actually used here (reference) ---') |
| print(' ' + truth[:200]) |
|
|