File size: 1,658 Bytes
6c71ff5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
#!/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()