#!/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()