polyedit-multiobjective-5seed / code /scripts /train_polyedit_multiobjective.py
promotion's picture
Upload code/scripts/train_polyedit_multiobjective.py with huggingface_hub
796b46a verified
Raw
History Blame Contribute Delete
11.5 kB
#!/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()