#!/usr/bin/env python """Train on seen MMP rule families and evaluate on disjoint held-out families.""" import argparse import hashlib import json from pathlib import Path import numpy as np from polyedit import policy, realenv from polyedit.rule_ood import build_rule_graphs from load_polyedit_bundle import load_verifier def split(canon): return "eval" if int(hashlib.sha1(canon.encode()).hexdigest(), 16) % 5 == 0 else "train" def evaluate(reqs, rollout, oracle): rows = [] for req in reqs: terminal = rollout(req) start = req.target.distance(req.source, oracle) final = min(req.target.distance(terminal, oracle), 1e3) rows.append({"source": req.source, "terminal": terminal, "success": realenv.hit(terminal, req), "regret": final, "direction": float(final < start)}) return {"n": len(rows), "success": float(np.mean([r["success"] for r in rows])), "regret": float(np.mean([r["regret"] for r in rows])), "direction": float(np.mean([r["direction"] for r in rows])), "tasks": rows} def main(): ap = argparse.ArgumentParser() ap.add_argument("--frozen", type=Path, default=Path("data/real/frozen.json")) ap.add_argument("--verifier_bundle", type=Path, required=True) ap.add_argument("--seed", type=int, required=True) ap.add_argument("--device", default="cuda") ap.add_argument("--epochs", type=int, default=60) ap.add_argument("--out", type=Path, required=True) ap.add_argument("--ckpt", type=Path, required=True) ap.add_argument("--wandb", action="store_true") args = ap.parse_args() records = json.loads(args.frozen.read_text())["polymers"] graphs, values, rule_meta = build_rule_graphs( records, ("Egc",), holdout_mod=2, holdout_bucket=0, min_support=3) if set(rule_meta["train_families"]) & set(rule_meta["test_families"]): raise SystemExit("rule-family leakage") train_polymers = [p for p in values if split(p) == "train"] eval_polymers = [p for p in values if split(p) == "eval"] common = dict(budget=3, rel_width=0.03, directional=True, min_delta=1.0) train_requests = ( realenv.build_requests(graphs["train"], values, ("Egc",), min_steps=1, sources=train_polymers, seed=0, max_requests=250, **common) + realenv.build_requests(graphs["train"], values, ("Egc",), min_steps=2, sources=train_polymers, seed=5, max_requests=125, **common)) tracks = { "iid_single": realenv.build_requests(graphs["train"], values, ("Egc",), min_steps=1, sources=eval_polymers, seed=1, max_requests=120, **common), "iid_multi": realenv.build_requests(graphs["train"], values, ("Egc",), min_steps=2, sources=eval_polymers, seed=2, max_requests=120, **common), "ood_single": realenv.build_requests(graphs["test"], values, ("Egc",), min_steps=1, sources=eval_polymers, seed=1, max_requests=120, **common), "ood_multi": realenv.build_requests(graphs["test"], values, ("Egc",), min_steps=2, sources=eval_polymers, seed=2, max_requests=120, **common), } if min(len(tracks["ood_single"]), len(tracks["ood_multi"])) < 30: raise SystemExit("fewer than 30 OOD tasks") run = None split_hash = hashlib.sha256("\n".join(rule_meta["test_families"]).encode()).hexdigest() if args.wandb: import wandb config = {"seed": args.seed, "epochs": args.epochs, "rule_split_hash": split_hash, "min_rule_support": 3, "holdout_fraction": 0.5, "train_families": len(rule_meta["train_families"]), "test_families": len(rule_meta["test_families"]), "tasks": {k: len(v) for k, v in tracks.items()}} run = wandb.init(entity="promotion-kim", project="polyedit", name=f"polyedit-real-rule-ood-s{args.seed}", config=config) log = (lambda row: run.log(row)) if run else None verifier, _ = load_verifier(args.verifier_bundle, args.device) polymers = (set(graphs["train"]) | set(graphs["test"]) | {p for graph in graphs.values() for row in graph.values() for p in row}) predictions = dict(zip(sorted(polymers), verifier.predict_many(sorted(polymers)))) verifiers = {"Egc": policy.CachedPredictor(verifier, predictions)} embeddings = policy.encode_polymers(sorted(polymers), device=args.device) steps = realenv.build_sft_steps(train_requests, graphs["train"]) features, masks = policy.step_features(steps, train_requests, embeddings, verifiers) import torch torch.manual_seed(args.seed) model = policy.make_policy(len(next(iter(embeddings.values()))) + 1) model = policy.train_sft(model, features, masks, epochs=args.epochs, device=args.device, seed=args.seed, log=log) oracle = realenv.real_oracle(values) voracle = lambda p, prop: verifiers[prop].predict(p) results = {} for track, reqs in tracks.items(): graph = graphs["test"] if track.startswith("ood") else graphs["train"] results[track] = { "sft": evaluate(reqs, lambda r, g=graph: realenv.policy_rollout( model, r, g, embeddings, verifiers, args.device), oracle), "random": evaluate(reqs, lambda r, g=graph: realenv.random_rollout( r.source, g, r.target, oracle, budget=r.budget, seed=args.seed), oracle), "greedy_verifier": evaluate(reqs, lambda r, g=graph: realenv.greedy_rollout( r.source, g, r.target, voracle, budget=r.budget), oracle), "greedy_oracle": evaluate(reqs, lambda r, g=graph: realenv.greedy_rollout( r.source, g, r.target, oracle, budget=r.budget), oracle), } if run: run.log({f"rule_ood/{track}/{method}/{metric}": row[metric] for method, row in results[track].items() for metric in ("success", "regret", "direction")}) output = {"seed": args.seed, "rule_split_hash": split_hash, "rule_meta": rule_meta, "n_train_requests": len(train_requests), "n_steps": len(steps), "results": results, "wandb_run": run.url if run else None} args.out.parent.mkdir(parents=True, exist_ok=True) args.out.write_text(json.dumps(output, indent=2) + "\n") args.ckpt.mkdir(parents=True, exist_ok=True) torch.save({"format_version": 1, "state_dict": {k: v.detach().cpu() for k, v in model.state_dict().items()}, "mu": torch.from_numpy(model.mu_), "sigma": torch.from_numpy(model.sigma_), "feat_dim": len(next(iter(embeddings.values()))) + 1, "polybert": policy.POLYBERT, "rule_split_hash": split_hash, "seed": args.seed}, args.ckpt / "sft_policy_bundle.pt") (args.ckpt / "config.json").write_text(json.dumps( {"seed": args.seed, "rule_split_hash": split_hash, "tasks": {k: len(v) for k, v in tracks.items()}}, indent=2) + "\n") if run: run.finish() print(f"seed={args.seed} tasks={ {k: len(v) for k, v in tracks.items()} }", flush=True) if __name__ == "__main__": main()