diff --git a/sgail-ppo-options-setobs2.py b/sgail-ppo-options-setobs2.py index e21ffdb..866930b 100644 --- a/sgail-ppo-options-setobs2.py +++ b/sgail-ppo-options-setobs2.py @@ -92,20 +92,16 @@ def training_function(config): expert_data = (torch.cat(d0), torch.cat(d1), torch.cat(d2), torch.cat(d3)) expert_data = Buffer(*expert_data) - run_folder = str(datetime.now()) - os.mkdir(os.path.join(DIR, run_folder)) - with open(os.path.join(DIR, run_folder, 'config.json'), 'w') as f: - json.dump(config, f, indent=4) - def callback(info): tune.report(gen_mean_reward_per_episode=info['gen/mean_reward_per_episode'], disc_mean_reward_per_episode=info['disc/mean_reward_per_episode'], - mean_episode_length=info['gen/mean_episode_length']) + mean_episode_length=info['gen/mean_episode_length'], + gen_collision_rate=info['gen/collision_rate']) # save model checkpoints ep = info['epoch'] + 1 if (ep % 25 == 0): - torch.save(info['policy'].state_dict(), os.path.join(DIR, run_folder, f'policy_epoch{ep}.pt')) + torch.save(info['policy'].state_dict(), f'policy_epoch{ep}.pt') value, policy = gail_ppo( env_fn=env_fn, @@ -164,7 +160,7 @@ analysis = tune.run( } ) -print('Best config: ', analysis.get_best_config(metric='gen_mean_reward_per_episode', mode='max')) +print('Best config: ', analysis.get_best_config(metric='gen_collision_rate', mode='min')) # %% # policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) diff --git a/src/safe_options/options.py b/src/safe_options/options.py index 662e6bb..7e160af 100644 --- a/src/safe_options/options.py +++ b/src/safe_options/options.py @@ -57,7 +57,8 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value gen_mean_reward_per_episode = generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0] logger.add_scalar('gen/mean_reward_per_episode', gen_mean_reward_per_episode, epoch) logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch) - logger.add_scalar('gen/collision_rate', (1. * collisions.any(-1)).mean(), epoch) + gen_collision_rate = (1. * collisions.any(-1)).mean() + logger.add_scalar('gen/collision_rate', gen_collision_rate, epoch) discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) if wasserstein: @@ -81,6 +82,7 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value 'policy': policy, 'gen/mean_episode_length': gen_mean_episode_length.item(), 'gen/mean_reward_per_episode': gen_mean_reward_per_episode.item(), + 'gen/collision_rate': gen_collision_rate.item(), 'disc/mean_reward_per_episode': disc_mean_reward_per_episode.item(), }) @@ -103,7 +105,8 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v gen_mean_reward_per_episode = generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0] logger.add_scalar('gen/mean_reward_per_episode', gen_mean_reward_per_episode, epoch) logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch) - logger.add_scalar('gen/collision_rate', (1. * collisions.any(-1)).mean(), epoch) + gen_collision_rate = (1. * collisions.any(-1)).mean() + logger.add_scalar('gen/collision_rate', gen_collision_rate, epoch) discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) if wasserstein: @@ -127,6 +130,7 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v 'policy': policy, 'gen/mean_episode_length': gen_mean_episode_length.item(), 'gen/mean_reward_per_episode': gen_mean_reward_per_episode.item(), + 'gen/collision_rate': gen_collision_rate.item(), 'disc/mean_reward_per_episode': disc_mean_reward_per_episode.item(), })