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