tiny-cube-value / code /train_value.py
briscoooe's picture
Value function training and state search
223c8ed verified
Raw
History Blame Contribute Delete
9.37 kB
"""
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()