File size: 4,378 Bytes
32c0c6c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
#!/usr/bin/env python
"""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']:      # never feed more than ctx
            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]
        # one candidate = one tactic: the model writes "A  B" (double space) for two steps,
        # and a truncated tactic may end in a dangling connective.
        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])