polyedit-rule-ood-5seed / code /scripts /train_polyedit_rule_ood.py
promotion's picture
Upload code/scripts/train_polyedit_rule_ood.py with huggingface_hub
3c5f059 verified
Raw
History Blame Contribute Delete
7.39 kB
#!/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()