From 71e3c5f816ab8fc66fdf99f7204ed834c6fda714 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Sat, 26 Feb 2022 14:46:17 +0100 Subject: [PATCH] Parametrize discriminator architecture --- sgail-ppo-options-setobs2.py | 11 ++++++++++- src/core/discriminator.py | 26 ++++++++++---------------- 2 files changed, 20 insertions(+), 17 deletions(-) diff --git a/sgail-ppo-options-setobs2.py b/sgail-ppo-options-setobs2.py index 866930b..09eca49 100644 --- a/sgail-ppo-options-setobs2.py +++ b/sgail-ppo-options-setobs2.py @@ -76,7 +76,12 @@ def training_function(config): value = SetValue() # config net architecture v_opt = torch.optim.Adam(value.parameters(), lr=config['value']['learning_rate']) - discriminator = DeepsetDiscriminator() # config net architecture + discriminator = DeepsetDiscriminator( + n_hidden_layers_element=config['discriminator']['n_hidden_layers_element'], + n_hidden_layers_global=config['discriminator']['n_hidden_layers_global'], + hidden_layer_size=config['discriminator']['hidden_layer_size'], + activation=config['discriminator']['activation'], + ) disc_opt = torch.optim.Adam(discriminator.parameters(), lr=config['discriminator']['learning_rate'], weight_decay=config['discriminator']['weight_decay']) expert_data = [ @@ -155,6 +160,10 @@ analysis = tune.run( 'learning_rate': 1e-3, #tune.grid_search([1e-3]), 'weight_decay': 1e-4, #tune.grid_search([1e-4]), 'iterations_per_epoch': 100, #tune.grid_search([100]), + 'n_hidden_layers_element': 3, + 'n_hidden_layers_global': 2, + 'hidden_layer_size': 10, + 'activation': torch.nn.Tanh, }, 'seed': 0, } diff --git a/src/core/discriminator.py b/src/core/discriminator.py index 8074c32..f4aaa4e 100644 --- a/src/core/discriminator.py +++ b/src/core/discriminator.py @@ -18,23 +18,17 @@ class Discriminator(nn.Module): class DeepsetDiscriminator(nn.Module): - def __init__(self): + def __init__(self, n_hidden_layers_element=3, n_hidden_layers_global=2, hidden_layer_size=10, activation=nn.Tanh): super().__init__() - self.elem = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(1), - ) - + + layers_elem = sum([[nn.LazyLinear(hidden_layer_size), + activation()] for _ in range(n_hidden_layers_element)], []) + self.elem = nn.Sequential(*layers_elem) + + layers_glob = sum([[nn.LazyLinear(hidden_layer_size), + activation()] for _ in range(n_hidden_layers_global)], []) + self.glob = nn.Sequential(*layers_glob, nn.LazyLinear(1)) + def forward(self, states, actions): actions = actions.unsqueeze(-2) actions = actions.expand(*actions.shape[:-2], states.shape[-2], actions.shape[-1])