#!/usr/bin/env python """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()