| 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, 'Pusher-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-Pusher-v5/weights.npz') |
|
|
| def main(): |
| parser = argparse.ArgumentParser(description='Enjoy Hybrid-HD-PPO on Pusher-v5') |
| parser.add_argument('--weights', default='hdppo-Pusher-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('Pusher-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() |
|
|