polyedit-multiobjective-5seed / load_polyedit_bundle.py
promotion's picture
Upload load_polyedit_bundle.py with huggingface_hub
6c71ff5 verified
Raw
History Blame Contribute Delete
1.66 kB
#!/usr/bin/env python
import argparse
from pathlib import Path
import torch
def load_policy(path, device="cpu"):
from polyedit.policy import make_policy
bundle = torch.load(path, map_location=device, weights_only=True)
model = make_policy(bundle["feat_dim"])
model.load_state_dict(bundle["state_dict"])
model.mu_ = bundle["mu"].cpu().numpy()
model.sigma_ = bundle["sigma"].cpu().numpy()
return model.to(device).eval(), bundle
def load_verifier(path, device="cpu"):
from transformers import AutoTokenizer
from polyedit.verifier_nn import FinetunedVerifier, _build_module
bundle = torch.load(path, map_location="cpu", weights_only=True)
model_name = bundle["model_name"]
module = _build_module(model_name, device)
module.load_state_dict(bundle["state_dict"])
tokenizer = AutoTokenizer.from_pretrained(model_name)
verifier = FinetunedVerifier(module, tokenizer, bundle["y_mean"], bundle["y_std"],
device, model_name)
return verifier, bundle
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--policy", type=Path, required=True)
ap.add_argument("--verifier", type=Path)
ap.add_argument("--device", default="cpu")
args = ap.parse_args()
policy, bundle = load_policy(args.policy, args.device)
print(f"policy loaded: feat_dim={bundle['feat_dim']}, params={sum(p.numel() for p in policy.parameters())}")
if args.verifier:
verifier, vb = load_verifier(args.verifier, args.device)
print(f"verifier loaded: property={vb['property']}, base={verifier.model_name}")
if __name__ == "__main__":
main()