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

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