Add learning rate schedule to SHAIL-PPO
This commit is contained in:
@@ -49,6 +49,7 @@ env_fn = lambda i: envs[i]
|
|||||||
|
|
||||||
policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n)
|
policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n)
|
||||||
pi_opt = torch.optim.Adam(policy.parameters(), lr=3e-4)
|
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()
|
value = SetValue()
|
||||||
v_opt = torch.optim.Adam(value.parameters(), lr=1e-3)
|
v_opt = torch.optim.Adam(value.parameters(), lr=1e-3)
|
||||||
@@ -85,6 +86,7 @@ value, policy = gail_ppo(
|
|||||||
pi_iters=100,
|
pi_iters=100,
|
||||||
logger=SummaryWriter(comment='sgail-ppo-options-setobs2'),
|
logger=SummaryWriter(comment='sgail-ppo-options-setobs2'),
|
||||||
callback=callback,
|
callback=callback,
|
||||||
|
lr_schedulers=[pi_lr_scheduler],
|
||||||
)
|
)
|
||||||
|
|
||||||
torch.save(policy.state_dict(), 'sgail-ppo-options-setobs2.pt')
|
torch.save(policy.state_dict(), 'sgail-ppo-options-setobs2.pt')
|
||||||
|
|||||||
@@ -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,
|
def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
|
||||||
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma,
|
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_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])
|
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:
|
if callback is not None:
|
||||||
callback(epoch, value, policy)
|
callback(epoch, value, policy)
|
||||||
|
|
||||||
|
for lr_scheduler in lr_schedulers:
|
||||||
|
lr_scheduler.step()
|
||||||
|
|
||||||
return value, policy
|
return value, policy
|
||||||
|
|
||||||
def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
|
def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
|
||||||
|
|||||||
Reference in New Issue
Block a user