File size: 9,365 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
"""
Trains a value function: cube state -> distance to solved.

## Why this and not another policy

Four attempts to make a sequence-predicting policy reach past ~8 moves all
failed (capacity, depth, chunked inference, training on its own rollouts), and
beam search over that policy raised depth 6 to 94% while leaving depth 15+ at
exactly zero. Beam search explores only what the policy already ranks highly; it
has no independent notion of which states are closer to solved, so where the
policy is wrong there is nothing better to find.

A value function supplies that missing compass. This is DeepCubeA's shape
(Agostinelli et al., Nature MI 2019), which solves the full cube with a network
of roughly this size by learning distance-to-solved and searching over it.

## Why this is cheaper here than in the paper

DeepCubeA has no expert, so it bootstraps the value function by approximate
value iteration -- expensive, and the bulk of its compute. We *do* have an
expert: Kociemba's solution length is the distance, and every training row
already carries it. So this is plain supervised learning on labels we generate
for free.

The labels are Kociemba two-phase lengths, which are upper bounds rather than
true optimal distances. That makes the heuristic inadmissible (A* over it is not
guaranteed optimal), which does not matter here: the goal is to find *a*
solution, not the shortest. The lengths were separately confirmed monotonic along
solution paths, so the signal is consistent.

## Classification, not regression

Distance is predicted as a distribution over 0..MAX_DIST rather than a scalar.
Regression to a mean is actively harmful on this task: the state space is
overwhelmingly distance 18-21, so a squared-error model collapses toward that
mode and loses exactly the near-solved discrimination the search depends on.

**The number to watch is mean absolute error.** Search needs the heuristic to
rank neighbouring states correctly; at MAE around 1 move that works, and at MAE
of several moves it cannot, whatever the loss curve says.
"""
import argparse, json, math, 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

MAX_DIST = 26


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

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

    def __getitem__(self, i):
        state, dist = self.rows[i]
        return torch.tensor(T.encode_state(state)), torch.tensor(dist)


def load_value_rows(path, limit=0):
    """Label is the solution length -- the distance -- regardless of how the row
    was generated, so augmented and canonical rows are equally usable."""
    rows = []
    with open(path) as fh:
        for line in fh:
            if limit and len(rows) >= limit:
                break
            r = json.loads(line)
            d = len(r["solution"].split())
            if 0 < d <= MAX_DIST:
                rows.append((r["state"], d))
    return rows


class ValueNet(torch.nn.Module):
    """Llama encoder over the 56-token state, mean-pooled, then a distance head."""

    def __init__(self, hidden, layers, heads):
        super().__init__()
        from transformers import LlamaConfig, LlamaModel
        cfg = LlamaConfig(
            vocab_size=T.VOCAB_SIZE, hidden_size=hidden, intermediate_size=hidden * 4,
            num_hidden_layers=layers, num_attention_heads=heads, num_key_value_heads=heads,
            max_position_embeddings=T.MAX_SEQ, pad_token_id=T.PAD,
        )
        self.encoder = LlamaModel(cfg)
        self.head = torch.nn.Linear(hidden, MAX_DIST + 1)

    def forward(self, ids):
        h = self.encoder(input_ids=ids).last_hidden_state.mean(dim=1)
        return self.head(h)


@torch.no_grad()
def evaluate(model, rows, device, batch_size=512):
    """Reports MAE and exact-match, plus MAE split by distance band.

    The split matters: an aggregate MAE can look fine while the near-solved band
    -- the only place the heuristic has to be sharp for search to make progress
    -- is useless.
    """
    model.eval()
    abs_err, exact, n = 0.0, 0, 0
    bands = defaultdict(lambda: [0.0, 0])
    for i in range(0, len(rows), batch_size):
        chunk = rows[i:i + batch_size]
        ids = torch.tensor([T.encode_state(s) for s, _ in chunk], device=device)
        true = torch.tensor([d for _, d in chunk], device=device)
        pred = model(ids).argmax(-1)
        err = (pred - true).abs().float()
        abs_err += err.sum().item()
        exact += (pred == true).sum().item()
        n += len(chunk)
        for e, t in zip(err.tolist(), true.tolist()):
            key = "1-8" if t <= 8 else ("9-14" if t <= 14 else "15+")
            bands[key][0] += e
            bands[key][1] += 1
    model.train()
    parts = " ".join(f"{k}:{bands[k][0]/bands[k][1]:.2f}" for k in ("1-8", "9-14", "15+") if bands[k][1])
    return abs_err / n, exact / n, parts


def main():
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("--data", required=True)
    p.add_argument("--val-data", default="")
    p.add_argument("--limit", type=int, default=0)
    p.add_argument("--hidden", type=int, default=512)
    p.add_argument("--layers", type=int, default=8)
    p.add_argument("--heads", type=int, default=8)
    p.add_argument("--batch-size", type=int, default=1024)
    p.add_argument("--lr", type=float, default=6e-4)
    p.add_argument("--max-steps", type=int, default=30000)
    p.add_argument("--warmup", type=int, default=200)
    p.add_argument("--eval-every", type=int, default=1000)
    p.add_argument("--eval-n", type=int, default=4096)
    p.add_argument("--out", default="checkpoints/value")
    p.add_argument("--hub-repo", default="")
    p.add_argument("--seed", type=int, default=0)
    args = p.parse_args()

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

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

    model = ValueNet(args.hidden, args.layers, args.heads).to(device)
    print(f"params {sum(q.numel() for q in model.parameters())/1e6:.1f}M", flush=True)

    loader = DataLoader(ValueDataset(rows), batch_size=args.batch_size, shuffle=True,
                        drop_last=True, num_workers=2)
    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(args.max_steps, 1), 1.0))))
    use_amp = device == "cuda"
    amp_dtype = torch.bfloat16 if (use_amp and torch.cuda.is_bf16_supported()) else torch.float16
    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(), 1e9
    while step < args.max_steps:
        for ids, dist in loader:
            ids, dist = ids.to(device), dist.to(device)
            with torch.autocast("cuda", dtype=amp_dtype, enabled=use_amp):
                loss = F.cross_entropy(model(ids), dist)
            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 % 100 == 0:
                print(f"step {step}/{args.max_steps} loss {loss.item():.4f} "
                      f"{step/(time.time()-t0):.1f} it/s", flush=True)
            if step % args.eval_every == 0 or step >= args.max_steps:
                mae, exact, parts = evaluate(model, val, device)
                print(f"  step {step}  MAE {mae:.3f} moves  exact {exact:.1%}  "
                      f"by true distance: {parts}", flush=True)
                torch.save({"state_dict": model.state_dict(),
                            "hidden": args.hidden, "layers": args.layers,
                            "heads": args.heads}, f"{args.out}/value.pt")
                if mae < best:
                    best = mae
                    torch.save({"state_dict": model.state_dict(),
                                "hidden": args.hidden, "layers": args.layers,
                                "heads": args.heads}, f"{args.out}/value-best.pt")
                if args.hub_repo:
                    try:
                        from huggingface_hub import HfApi
                        HfApi().upload_file(path_or_fileobj=f"{args.out}/value.pt",
                                            path_in_repo="value.pt", repo_id=args.hub_repo)
                    except Exception as e:
                        print(f"  hub push failed (continuing): {e}", flush=True)
            if step >= args.max_steps:
                break
    print(f"done. best MAE {best:.3f} moves. weights in {args.out}", flush=True)


if __name__ == "__main__":
    main()