Add learning rate schedule to SHAIL-PPO

This commit is contained in:
ebuehrle
2022-02-23 14:22:02 +01:00
parent 68b066ec53
commit 91f88983b0
2 changed files with 6 additions and 1 deletions

View File

@@ -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')

View File

@@ -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):