| |
| """Preference-conditioned Egc/Egb SFT with frozen-DFT Hypervolume and MIP evaluation.""" |
| import argparse |
| import hashlib |
| import json |
| from pathlib import Path |
|
|
| import numpy as np |
|
|
| from polyedit import policy, realenv, verifier_nn |
| from polyedit.multiobjective import (hypervolume_2d, induced_graph, mip, |
| preference_grid, reward_vector) |
|
|
|
|
| PROPS = ("Egc", "Egb") |
| DIRECTIONS = ((1, 1), (1, -1), (-1, 1), (-1, -1)) |
| POLYBERT_REVISION = "7bf9ed32ac54dea5bc163cf90100728b49341750" |
|
|
|
|
| def is_eval(poly): |
| return int(hashlib.sha1(poly.encode()).hexdigest(), 16) % 2 == 0 |
|
|
|
|
| def candidates(state, graph): |
| return list(graph.get(state, ())) + [state] |
|
|
|
|
| def feature_rows(polys, directions, weight, emb, predictions, means, scales): |
| pref = np.asarray(weight) * np.asarray(directions) |
| rows = [] |
| for poly in polys: |
| z = np.asarray([(predictions[p][poly] - means[p]) / scales[p] for p in PROPS]) |
| rows.append(np.concatenate([emb[poly], z, pref, pref * z])) |
| return np.asarray(rows, dtype=np.float32) |
|
|
|
|
| def score_policy(model, polys, directions, weight, emb, predictions, means, scales, device): |
| import torch |
|
|
| x = feature_rows(polys, directions, weight, emb, predictions, means, scales) |
| x = (x - model.mu_) / model.sigma_ |
| with torch.no_grad(): |
| return model(torch.from_numpy(x).to(device)).squeeze(-1).cpu().numpy() |
|
|
|
|
| def greedy(source, graph, budget, utility): |
| state = source |
| for _ in range(budget): |
| options = candidates(state, graph) |
| best = max(options, key=lambda p: (utility(p), p)) |
| if best == state: |
| break |
| state = best |
| return state |
|
|
|
|
| def random_rollout(source, graph, budget, seed): |
| rng, state = np.random.default_rng(seed), source |
| for _ in range(budget): |
| options = graph.get(state, ()) |
| if not options: |
| break |
| state = options[int(rng.integers(len(options)))] |
| return state |
|
|
|
|
| def task_sources(graph): |
| return [p for p in sorted(graph) if graph[p]] |
|
|
|
|
| def build_sft_data(graph, values, emb, predictions, means, scales, budget, weights): |
| features, masks = [], [] |
| for source in task_sources(graph): |
| reachable = {source: ()} | realenv.reachable_all(source, graph, budget=budget) |
| for directions in DIRECTIONS: |
| rewards = {p: reward_vector(p, directions, values, means, scales) for p in reachable} |
| for weight in weights: |
| terminal = max(reachable, key=lambda p: (float(np.dot(weight, rewards[p])), p)) |
| state = source |
| for nxt in tuple(reachable[terminal]) + (terminal,): |
| options = candidates(state, graph) |
| x = feature_rows(options, directions, weight, emb, predictions, means, scales) |
| mask = np.asarray([p == nxt for p in options], dtype=bool) |
| features.append(x); masks.append(mask) |
| if nxt == state: |
| break |
| state = nxt |
| return features, masks |
|
|
|
|
| def evaluate(model, graph, values, emb, predictions, means, scales, budget, weights, seed, device): |
| methods = ("sft", "random", "greedy_verifier", "oracle_grid") |
| predicted_values = {q: {p: predictions[p][q] for p in PROPS} for q in graph} |
| rows = {m: [] for m in methods} |
| for ti, source in enumerate(task_sources(graph)): |
| reachable = {source: ()} | realenv.reachable_all(source, graph, budget=budget) |
| for directions in DIRECTIONS: |
| true_reward = {p: reward_vector(p, directions, values, means, scales) for p in reachable} |
| outputs = {m: [] for m in methods} |
| for wi, weight in enumerate(weights): |
| true_u = lambda p, w=weight: float(np.dot(w, true_reward[p])) |
| pred_u = lambda p, w=weight: float(np.dot( |
| w, reward_vector(p, directions, predicted_values, means, scales))) |
| outputs["sft"].append(greedy( |
| source, graph, budget, |
| lambda p, w=weight: float(score_policy( |
| model, [p], directions, w, emb, predictions, means, scales, device)[0]))) |
| outputs["random"].append(random_rollout( |
| source, graph, budget, seed + 1009 * ti + 37 * wi)) |
| outputs["greedy_verifier"].append(greedy(source, graph, budget, pred_u)) |
| outputs["oracle_grid"].append(max(reachable, key=lambda p: (true_u(p), p))) |
|
|
| oracle_u = [float(np.dot(w, true_reward[p])) |
| for w, p in zip(weights, outputs["oracle_grid"])] |
| for method in methods: |
| rewards = [true_reward[p] for p in outputs[method]] |
| utility = [float(np.dot(w, r)) for w, r in zip(weights, rewards)] |
| rows[method].append({ |
| "source": source, "directions": list(directions), |
| "hypervolume": hypervolume_2d(rewards), "mip": mip(weights, rewards), |
| "preference_regret": float(np.mean(np.asarray(oracle_u) - utility)), |
| "near_optimal": float(np.mean(np.asarray(oracle_u) - utility <= 0.05)), |
| "unique": len(set(outputs[method])), |
| }) |
| return {method: {"n_tasks": len(items), |
| **{metric: float(np.mean([x[metric] for x in items])) |
| for metric in ("hypervolume", "mip", "preference_regret", "near_optimal", "unique")}, |
| "tasks": items} for method, items in rows.items()} |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--frozen", type=Path, default=Path("data/real/frozen.json")) |
| ap.add_argument("--polybert_path", required=True) |
| ap.add_argument("--seed", type=int, required=True) |
| ap.add_argument("--device", default="cuda") |
| ap.add_argument("--budget", type=int, default=3) |
| ap.add_argument("--epochs", type=int, default=60) |
| ap.add_argument("--ft_epochs", type=int, default=8) |
| 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() |
|
|
| frozen = json.loads(args.frozen.read_text()) |
| graph, values, _ = realenv.build_graph(frozen["polymers"], PROPS) |
| train_nodes = [p for p in values if not is_eval(p)] |
| eval_nodes = [p for p in values if is_eval(p)] |
| train_graph, eval_graph = induced_graph(graph, train_nodes), induced_graph(graph, eval_nodes) |
| weights = preference_grid(11) |
| means = {p: float(np.mean([values[x][p] for x in train_nodes])) for p in PROPS} |
| scales = {p: float(np.std([values[x][p] for x in train_nodes])) or 1.0 for p in PROPS} |
| split_hash = hashlib.sha256("\n".join(sorted(eval_nodes)).encode()).hexdigest() |
| if set(train_graph) & set(eval_graph) or len(task_sources(eval_graph)) * len(DIRECTIONS) < 50: |
| raise SystemExit("invalid split or fewer than 50 evaluation tasks") |
|
|
| run = None |
| if args.wandb: |
| import wandb |
| run = wandb.init(entity="promotion-kim", project="polyedit", |
| name=f"polyedit-multiobjective-s{args.seed}", config={ |
| "seed": args.seed, "budget": args.budget, "epochs": args.epochs, |
| "ft_epochs": args.ft_epochs, "split_hash": split_hash, |
| "preference_grid": len(weights), "directions": len(DIRECTIONS), |
| "polybert_revision": POLYBERT_REVISION}) |
| log = (lambda row: run.log(row)) if run else (lambda row: None) |
|
|
| verifiers, val_rmse = {}, {} |
| for prop in PROPS: |
| ver, rmse = verifier_nn.finetune( |
| train_nodes, [values[x][prop] for x in train_nodes], |
| eval_nodes, [values[x][prop] for x in eval_nodes], model_name=args.polybert_path, |
| epochs=args.ft_epochs, device=args.device, seed=args.seed, |
| log=lambda row, p=prop: log({f"verifier/{p}/{k.split('/')[-1]}": v |
| for k, v in row.items()})) |
| ver.model_name = policy.POLYBERT |
| verifiers[prop], val_rmse[prop] = ver, rmse |
|
|
| polymers = sorted(set(train_graph) | set(eval_graph)) |
| emb = policy.encode_polymers(polymers, device=args.device, model_name=args.polybert_path) |
| predictions = {p: dict(zip(polymers, verifiers[p].predict_many(polymers))) for p in PROPS} |
| features, masks = build_sft_data( |
| train_graph, values, emb, predictions, means, scales, args.budget, weights) |
|
|
| import torch |
| torch.manual_seed(args.seed) |
| model = policy.make_policy(features[0].shape[1]) |
| model = policy.train_sft(model, features, masks, epochs=args.epochs, device=args.device, |
| seed=args.seed, log=log) |
| results = evaluate(model, eval_graph, values, emb, predictions, means, scales, |
| args.budget, weights, args.seed, args.device) |
| for method, row in results.items(): |
| log({f"multiobjective/{method}/{metric}": row[metric] |
| for metric in ("hypervolume", "mip", "preference_regret", "near_optimal", "unique")}) |
|
|
| metadata = {"seed": args.seed, "data_sha256": hashlib.sha256(args.frozen.read_bytes()).hexdigest(), |
| "split_hash": split_hash, "train_labeled": len(train_nodes), |
| "eval_labeled": len(eval_nodes), "train_graph_sources": len(task_sources(train_graph)), |
| "eval_graph_sources": len(task_sources(eval_graph)), |
| "n_eval_tasks": len(task_sources(eval_graph)) * len(DIRECTIONS), |
| "n_preference_decisions": len(task_sources(eval_graph)) * len(DIRECTIONS) * len(weights), |
| "n_train_steps": len(features), "weights": weights, "directions": DIRECTIONS, |
| "reward_normalization": {"means": means, "scales": scales, "clip_sd": 3.0}, |
| "val_rmse": val_rmse, "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(metadata, 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": features[0].shape[1], "polybert": policy.POLYBERT, |
| "polybert_revision": POLYBERT_REVISION, "split_hash": split_hash, |
| "seed": args.seed}, args.ckpt / "preference_policy_bundle.pt") |
| for prop, ver in verifiers.items(): |
| torch.save({"format_version": 1, "property": prop, "model_name": policy.POLYBERT, |
| "model_revision": POLYBERT_REVISION, |
| "state_dict": {k: v.detach().cpu() for k, v in ver.module.state_dict().items()}, |
| "y_mean": ver.y_mean, "y_std": ver.y_std}, |
| args.ckpt / f"verifier_{prop}_bundle.pt") |
| (args.ckpt / "config.json").write_text(json.dumps( |
| {k: v for k, v in metadata.items() if k not in ("results",)}, indent=2) + "\n") |
| if run: |
| run.finish() |
| print(f"seed={args.seed} eval_tasks={metadata['n_eval_tasks']} metrics=" |
| f"{ {m: (round(r['hypervolume'], 4), round(r['mip'], 4)) for m, r in results.items()} }", |
| flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|