File size: 10,794 Bytes
223c8ed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
"""
Trains the tiny cube-solving model.

Architecture is a stock LlamaConfig rather than a bespoke nn.Module. The model is
identical either way, but a standard architecture loads with plain
`transformers` (no trust_remote_code), pushes to the Hub cleanly, and stays
compatible with the wider tooling if it is ever wanted. That costs a config
object instead of a class.

## The metric is solve rate, not loss

Token accuracy and validation loss are both misleading here. A cube has
astronomically many valid solutions and Kociemba emits one of them, so a model
that produces a *different* valid solve scores badly on token match and
perfectly on the only thing that matters. The harness scores by applying the
moves and asking the engine whether the cube ended solved, so evaluation here
does the same. Watch `solve_rate`; loss is only useful for spotting divergence.
"""
import argparse, json, math, os, random, sys, time
from collections import defaultdict
from pathlib import Path

import torch
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader

sys.path.insert(0, str(Path(__file__).parent))
import cube_tokenizer as T
from gen_data import SOLVED, apply_sequence


class CubeDataset(Dataset):
    def __init__(self, rows):
        self.rows = rows

    def __len__(self):
        return len(self.rows)

    def __getitem__(self, i):
        state, solution = self.rows[i]
        ids, labels = T.encode_pair(state, solution)
        return torch.tensor(ids), torch.tensor(labels)


def collate(batch):
    n = max(len(x[0]) for x in batch)
    ids = torch.full((len(batch), n), T.PAD, dtype=torch.long)
    labels = torch.full((len(batch), n), -100, dtype=torch.long)
    mask = torch.zeros((len(batch), n), dtype=torch.long)
    for i, (a, b) in enumerate(batch):
        ids[i, : len(a)] = a
        labels[i, : len(b)] = b
        mask[i, : len(a)] = 1
    return ids, labels, mask


def load_jsonl(path, limit=0):
    rows = []
    with open(path) as fh:
        for line in fh:
            if limit and len(rows) >= limit:
                break
            r = json.loads(line)
            moves = r["solution"].split()
            if len(moves) > T.MAX_SOLUTION:
                continue
            rows.append((r["state"], moves))
    return rows


@torch.no_grad()
def solve_rate(model, rows, device, max_new=T.MAX_SOLUTION + 1, batch_size=256):
    """Greedy-decodes each state and reports the fraction that solve, plus a
    breakdown by solution length.

    The breakdown is not decoration. A single aggregate over this holdout is
    saturated by construction: the set is ~15% near-solved states and ~85%
    fully-mixed ones, so a model whose reach stops at 8 moves cannot score above
    ~15% however well it trains. Watching only the aggregate showed a flat 11-12%
    for 13,000 steps while the model was in fact going from 52% to 72% on
    seven-move solves -- the progress was real and entirely invisible.

    Bucketed by label length rather than by the recorded scramble depth, since
    depth stops tracking difficulty past ~15 moves (every deep scramble is ~20
    moves from solved) and length is what actually governs whether the model can
    do it.
    """
    model.eval()
    solved = 0
    buckets = defaultdict(lambda: [0, 0])
    for start in range(0, len(rows), batch_size):
        chunk = rows[start : start + batch_size]
        prompts = torch.tensor([T.encode_state(s) for s, _ in chunk], device=device)
        out = prompts
        finished = torch.zeros(len(chunk), dtype=torch.bool, device=device)
        for _ in range(max_new):
            logits = model(input_ids=out).logits[:, -1, :]
            nxt = logits.argmax(-1)
            nxt[finished] = T.PAD
            finished |= nxt == T.EOS
            out = torch.cat([out, nxt[:, None]], dim=1)
            if finished.all():
                break
        for row, (state, label) in zip(out[:, prompts.shape[1] :].tolist(), chunk):
            moves = T.decode_solution(row)
            ok = bool(moves) and apply_sequence(state, moves) == SOLVED
            solved += ok
            key = "1-8" if len(label) <= 8 else ("9-14" if len(label) <= 14 else "15+")
            buckets[key][0] += ok
            buckets[key][1] += 1
    model.train()
    parts = " ".join(f"{k}:{buckets[k][0]}/{buckets[k][1]}"
                     for k in ("1-8", "9-14", "15+") if buckets[k][1])
    return solved / max(len(rows), 1), parts


def build_model(args):
    from transformers import LlamaConfig, LlamaForCausalLM

    cfg = LlamaConfig(
        vocab_size=T.VOCAB_SIZE,
        hidden_size=args.hidden,
        intermediate_size=args.hidden * 4,
        num_hidden_layers=args.layers,
        num_attention_heads=args.heads,
        num_key_value_heads=args.heads,
        max_position_embeddings=T.MAX_SEQ,
        rms_norm_eps=1e-5,
        pad_token_id=T.PAD,
        bos_token_id=T.BOS,
        eos_token_id=T.EOS,
        tie_word_embeddings=True,
    )
    return LlamaForCausalLM(cfg)


def main():
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("--data", required=True)
    p.add_argument("--val-data", default="", help="Held-out set; defaults to a slice of --data.")
    p.add_argument("--limit", type=int, default=0)
    p.add_argument("--hidden", type=int, default=256)
    p.add_argument("--layers", type=int, default=6)
    p.add_argument("--heads", type=int, default=8)
    p.add_argument("--batch-size", type=int, default=256)
    p.add_argument("--lr", type=float, default=3e-4)
    p.add_argument("--epochs", type=int, default=1)
    p.add_argument("--max-steps", type=int, default=0)
    p.add_argument("--warmup", type=int, default=200)
    p.add_argument("--eval-every", type=int, default=500)
    p.add_argument("--eval-n", type=int, default=256)
    p.add_argument("--out", default="checkpoints/cube")
    p.add_argument("--hub-repo", default="", help="Push checkpoints here (e.g. user/tiny-cube).")
    p.add_argument("--seed", type=int, default=0)
    p.add_argument("--resume", default="",
                   help="Checkpoint to continue from (local dir or Hub repo id).")
    args = p.parse_args()

    torch.manual_seed(args.seed)
    random.seed(args.seed)
    device = "cuda" if torch.cuda.is_available() else "cpu"

    rows = load_jsonl(args.data, args.limit)
    if args.val_data:
        val = load_jsonl(args.val_data, args.eval_n)
    else:
        split = max(len(rows) - args.eval_n, 1)
        rows, val = rows[:split], rows[split:]
    print(f"train {len(rows)} | val {len(val)} | device {device}", flush=True)

    # Resuming matters on a preemptible box: a reclaimed instance otherwise
    # restarts a multi-hour run from random init. Weights come back from the Hub,
    # which is why checkpoints are pushed there rather than kept only on disk.
    # Optimizer state is not restored -- only the weights -- so the LR schedule
    # restarts; that costs a little progress but keeps the checkpoint portable.
    if args.resume:
        from transformers import LlamaForCausalLM
        model = LlamaForCausalLM.from_pretrained(args.resume).to(device)
        print(f"resumed from {args.resume}", flush=True)
    else:
        model = build_model(args).to(device)
    n_params = sum(p.numel() for p in model.parameters())
    print(f"params {n_params/1e6:.1f}M", flush=True)

    loader = DataLoader(CubeDataset(rows), batch_size=args.batch_size, shuffle=True,
                        collate_fn=collate, drop_last=True, num_workers=2)
    steps = args.max_steps or len(loader) * args.epochs
    opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.1, betas=(0.9, 0.95))
    sched = torch.optim.lr_scheduler.LambdaLR(
        opt, lambda s: min((s + 1) / max(args.warmup, 1), 1.0)
        * 0.5 * (1 + math.cos(math.pi * min(s / max(steps, 1), 1.0))))

    # bf16 needs Ampere or newer. Falling back to fp16 rather than assuming, so a
    # cheaper pre-Ampere box (T4, V100, P100) trains correctly instead of silently
    # producing garbage or refusing to start.
    use_amp = device == "cuda"
    amp_dtype = torch.bfloat16
    if use_amp and not torch.cuda.is_bf16_supported():
        amp_dtype = torch.float16
        print("bf16 unsupported on this GPU; using fp16", flush=True)
    scaler = torch.amp.GradScaler("cuda", enabled=use_amp and amp_dtype is torch.float16)
    Path(args.out).mkdir(parents=True, exist_ok=True)
    step, t0, best = 0, time.time(), -1.0
    done = False

    while not done:
        for ids, labels, mask in loader:
            ids, labels, mask = ids.to(device), labels.to(device), mask.to(device)
            with torch.autocast("cuda", dtype=amp_dtype, enabled=use_amp):
                loss = model(input_ids=ids, attention_mask=mask, labels=labels).loss
            scaler.scale(loss).backward()
            scaler.unscale_(opt)
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            scaler.step(opt)
            scaler.update()
            sched.step()
            opt.zero_grad(set_to_none=True)
            step += 1

            if step % 50 == 0:
                print(f"step {step}/{steps} loss {loss.item():.4f} "
                      f"lr {sched.get_last_lr()[0]:.2e} {step/(time.time()-t0):.1f} it/s", flush=True)

            if step % args.eval_every == 0 or step == steps:
                rate, breakdown = solve_rate(model, val, device)
                print(f"  step {step}  SOLVE RATE {rate:.1%}  ({len(val)} held out)  "
                      f"by solution length: {breakdown}", flush=True)
                # Always keep the latest weights, and additionally keep the best.
                # Saving only on improvement silently threw away 22,000 steps on the
                # first real run: the holdout metric is saturated (see solve_rate),
                # so it peaked mid-run and every later checkpoint -- including the
                # final one -- was discarded. A metric that cannot distinguish two
                # models must not be the thing that chooses between them.
                model.save_pretrained(args.out)
                if rate >= best:
                    best = rate
                    model.save_pretrained(f"{args.out}-best")
                if args.hub_repo:
                    try:
                        model.push_to_hub(args.hub_repo, commit_message=f"step {step} solve {rate:.3f}")
                    except Exception as e:
                        print(f"  hub push failed (continuing): {e}", flush=True)

            if step >= steps:
                done = True
                break

    print(f"done. best solve rate {best:.1%}. final weights in {args.out}, "
          f"best-scoring in {args.out}-best", flush=True)


if __name__ == "__main__":
    main()