""" Scores the value function + state search on the canonical benchmark scrambles. Separate from eval_canonical.py because this is a different system: that one decodes a policy (greedy, beam, or chunked), this one searches the state space with a learned distance estimate as the heuristic and the engine as the transition function. The policy is not involved at all. Run on a GPU. Each solve evaluates `beam * 18` child states per step for up to `--max-depth` steps -- tens of thousands of network calls, which is nothing on a GPU and hopeless on a few CPU cores. python eval_value_search.py --value --beam 64 --limit-per-depth 25 """ import argparse, json, os, sys, time from collections import defaultdict from pathlib import Path sys.path.insert(0, str(Path(__file__).parent)) from gen_data import SOLVED, apply_sequence FIXTURE = Path(__file__).parent / "fixtures" / "canonical_states.json" def main(): p = argparse.ArgumentParser(description=__doc__) p.add_argument("--value", required=True, help="Path to a value checkpoint, or a Hub repo id to fetch it " "from. Either flavour works -- supervised (train_value.py) " "or value iteration (train_value_iteration.py); search.py " "detects which from the checkpoint.") p.add_argument("--value-file", default="value.pt", help="Filename within the Hub repo, when --value is a repo id.") p.add_argument("--beam", default="64", help="Beam width, or a comma-separated list to sweep. A sweep " "reuses one box for every width, which matters because " "the deep buckets are where width decides the outcome and " "a single guess at it wastes the whole run.") p.add_argument("--lam", default="0.0", help="Weight on path length: states are ranked by h + lam*g. " "0 is pure greedy on the heuristic (fastest, and longest " "solutions), 1.0 is standard A*. Accepts a comma-separated " "list to sweep.") p.add_argument("--max-depth", type=int, default=26) p.add_argument("--limit-per-depth", type=int, default=25) p.add_argument("--depths", default="") args = p.parse_args() import torch from search import load_value_model, solve_value_beam path = args.value if not os.path.exists(path): from huggingface_hub import hf_hub_download path = hf_hub_download(args.value, args.value_file, token=os.environ.get("HF_TOKEN")) device = "cuda" if torch.cuda.is_available() else "cpu" model = load_value_model(path, device) print(f"value model ({model.value_kind}) on {device} | beam {args.beam} | " f"lam {args.lam}", flush=True) states = json.loads(FIXTURE.read_text()) if args.depths: keep = {int(d) for d in args.depths.split(",")} states = [s for s in states if s["depth"] in keep] if args.limit_per_depth: seen = defaultdict(int) kept = [] for s in states: if seen[s["depth"]] < args.limit_per_depth: kept.append(s) seen[s["depth"]] += 1 states = kept grid = [(b, l) for b in [int(x) for x in str(args.beam).split(",")] for l in [float(x) for x in str(args.lam).split(",")]] for beam, lam in grid: print(f"\n=== beam {beam} lam {lam} ===", flush=True) by_depth = defaultdict(lambda: [0, 0, 0, 0.0]) print(f"{'depth':>6} {'solved':>8} {'n':>4} {'rate':>7} {'moves':>7} {'s/solve':>8}", flush=True) for s_ in states: t0 = time.time() moves = solve_value_beam(model, s_["facelets"], device, beam=beam, max_depth=args.max_depth, lam=lam) ok = bool(moves) and apply_sequence(s_["facelets"], moves) == SOLVED b = by_depth[s_["depth"]] b[0] += ok b[1] += 1 b[2] += len(moves) if ok else 0 b[3] += time.time() - t0 # Printed per depth as each depth's last scramble lands, rather than # only at the end: a wide beam on the deep buckets can run for a long # time, and a box that is killed or times out mid-sweep should still # have reported every bucket it did finish. if b[1] == sum(1 for x in states if x["depth"] == s_["depth"]): ok_, n_, mv_, secs_ = b print(f"{s_['depth']:>6} {ok_:>8} {n_:>4} {ok_/n_:>6.1%} " f"{mv_/max(ok_,1):>7.1f} {secs_/n_:>8.1f}", flush=True) tot_ok = sum(v[0] for v in by_depth.values()) tot_n = sum(v[1] for v in by_depth.values()) print(f"{'ALL':>6} {tot_ok:>8} {tot_n:>4} {tot_ok/tot_n:>6.1%}", flush=True) if __name__ == "__main__": main()