From 91f88983b01521fdc48a30d0b8ccee8da7c6e7b6 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Wed, 23 Feb 2022 14:22:02 +0100 Subject: [PATCH] Add learning rate schedule to SHAIL-PPO --- .../etienne/trpo/experiments/sgail-ppo-options-setobs2.py | 2 ++ src/safe_options/options.py | 5 ++++- 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py index a7b01e7..11bb889 100644 --- a/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py @@ -49,6 +49,7 @@ env_fn = lambda i: envs[i] policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) pi_opt = torch.optim.Adam(policy.parameters(), lr=3e-4) +pi_lr_scheduler = torch.optim.lr_scheduler.StepLR(pi_opt, step_size=50, gamma=0.2) value = SetValue() v_opt = torch.optim.Adam(value.parameters(), lr=1e-3) @@ -85,6 +86,7 @@ value, policy = gail_ppo( pi_iters=100, logger=SummaryWriter(comment='sgail-ppo-options-setobs2'), callback=callback, + lr_schedulers=[pi_lr_scheduler], ) torch.save(policy.state_dict(), 'sgail-ppo-options-setobs2.pt') diff --git a/src/safe_options/options.py b/src/safe_options/options.py index 02460af..1042735 100644 --- a/src/safe_options/options.py +++ b/src/safe_options/options.py @@ -77,7 +77,7 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): + gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None, lr_schedulers=[]): logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) @@ -109,6 +109,9 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v if callback is not None: callback(epoch, value, policy) + for lr_scheduler in lr_schedulers: + lr_scheduler.step() + return value, policy def rollout(env_fn, policy, n_episodes, max_steps_per_episode):