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()
|