| |
| 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() |
|
|
|
|