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