File size: 3,588 Bytes
0c0f8ca
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
import argparse
import os
import gymnasium as gym
import numpy as np
import torch
import train_hdppo as m

def _warmup_policy(encoder, actor, env_id, env_kwargs, n_steps=200):
    env = gym.make(env_id, **env_kwargs)
    obs, _ = env.reset(seed=0)
    for i in range(n_steps):
        Hr, Hi = encoder.encode(obs)
        action = actor.greedy_action_np(Hr, Hi)
        obs, _, term, trunc, _ = env.step(action)
        if term or trunc:
            obs, _ = env.reset(seed=i + 1)
    env.close()

def load_policy_from_checkpoint(path, seed=42, warmup=True):
    data = np.load(path)
    cfg = dict(m.CONFIG)
    cfg['D'] = int(data['D'])
    cfg['beta'] = float(data['beta_base'])
    cfg['fpe_phi_init'] = data['fpe_phi']
    if 'feat_lo' in data:
        cfg['feat_lo'] = data['feat_lo'].tolist()
        cfg['feat_hi'] = data['feat_hi'].tolist()
    torch.manual_seed(seed)
    np.random.seed(seed)
    encoder = m.make_encoder(cfg, seed)
    actor = m.HDLinearActor(cfg['D'], cfg['action_dim'], cfg['log_std_init'], float(data['action_low']) if 'action_low' in data else cfg['action_low'], float(data['action_high']) if 'action_high' in data else cfg['action_high'])
    with torch.no_grad():
        actor.W_re.data.copy_(torch.from_numpy(np.asarray(data['W_actor_re'], dtype=np.float32)))
        actor.W_im.data.copy_(torch.from_numpy(np.asarray(data['W_actor_im'], dtype=np.float32)))
        if 'log_std' in data:
            actor.log_std.data.copy_(torch.from_numpy(np.asarray(data['log_std'], dtype=np.float32)))
    if cfg.get('adaptive_beta', True):
        encoder.link_torch_actor(actor)
    actor.eval()
    if warmup:
        _warmup_policy(encoder, actor, 'InvertedPendulum-v5', dict(cfg.get('env_kwargs', {})))
    return (encoder, actor, cfg)

def resolve_weights(weights_arg):
    if os.path.isfile(weights_arg):
        return weights_arg
    try:
        from huggingface_hub import hf_hub_download
    except ImportError as exc:
        raise SystemExit('Install huggingface_hub to load remote checkpoints: pip install huggingface_hub') from exc
    return hf_hub_download(repo_id=weights_arg, filename='hdppo-InvertedPendulum-v5/weights.npz')

def main():
    parser = argparse.ArgumentParser(description='Enjoy Hybrid-HD-PPO on InvertedPendulum-v5')
    parser.add_argument('--weights', default='hdppo-InvertedPendulum-v5/weights.npz')
    parser.add_argument('--episodes', type=int, default=5)
    parser.add_argument('--seed', type=int, default=10000)
    parser.add_argument('--render', action='store_true')
    args = parser.parse_args()
    weights = resolve_weights(args.weights)
    encoder, actor, cfg = load_policy_from_checkpoint(weights, seed=args.seed)
    render_mode = 'human' if args.render else None
    env_kwargs = dict(cfg.get('env_kwargs', {}))
    if render_mode is not None:
        env_kwargs['render_mode'] = render_mode
    env = gym.make('InvertedPendulum-v5', **env_kwargs)
    rewards = []
    for ep in range(args.episodes):
        obs, _ = env.reset(seed=args.seed + ep)
        done = False
        ep_r = 0.0
        while not done:
            Hr, Hi = encoder.encode(obs)
            action = actor.greedy_action_np(Hr, Hi)
            obs, reward, term, trunc, _ = env.step(action)
            ep_r += float(reward)
            done = term or trunc
        rewards.append(ep_r)
        print(f'episode {ep + 1}: return={ep_r:.1f}')
    env.close()
    arr = np.asarray(rewards, dtype=np.float64)
    print(f'mean={arr.mean():.2f}  std={(arr.std(ddof=1) if len(arr) > 1 else 0.0):.2f}')
if __name__ == '__main__':
    main()