Parametrize discriminator architecture

This commit is contained in:
ebuehrle
2022-02-26 14:46:17 +01:00
parent 35e6fb299c
commit 71e3c5f816
2 changed files with 20 additions and 17 deletions

View File

@@ -76,7 +76,12 @@ def training_function(config):
value = SetValue() # config net architecture value = SetValue() # config net architecture
v_opt = torch.optim.Adam(value.parameters(), lr=config['value']['learning_rate']) 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']) disc_opt = torch.optim.Adam(discriminator.parameters(), lr=config['discriminator']['learning_rate'], weight_decay=config['discriminator']['weight_decay'])
expert_data = [ expert_data = [
@@ -155,6 +160,10 @@ analysis = tune.run(
'learning_rate': 1e-3, #tune.grid_search([1e-3]), 'learning_rate': 1e-3, #tune.grid_search([1e-3]),
'weight_decay': 1e-4, #tune.grid_search([1e-4]), 'weight_decay': 1e-4, #tune.grid_search([1e-4]),
'iterations_per_epoch': 100, #tune.grid_search([100]), '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, 'seed': 0,
} }

View File

@@ -18,23 +18,17 @@ class Discriminator(nn.Module):
class DeepsetDiscriminator(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__() super().__init__()
self.elem = nn.Sequential(
nn.LazyLinear(10), layers_elem = sum([[nn.LazyLinear(hidden_layer_size),
nn.Tanh(), activation()] for _ in range(n_hidden_layers_element)], [])
nn.LazyLinear(10), self.elem = nn.Sequential(*layers_elem)
nn.Tanh(),
nn.LazyLinear(10), layers_glob = sum([[nn.LazyLinear(hidden_layer_size),
) activation()] for _ in range(n_hidden_layers_global)], [])
self.glob = nn.Sequential( self.glob = nn.Sequential(*layers_glob, nn.LazyLinear(1))
nn.LazyLinear(10),
nn.Tanh(),
nn.LazyLinear(10),
nn.Tanh(),
nn.LazyLinear(1),
)
def forward(self, states, actions): def forward(self, states, actions):
actions = actions.unsqueeze(-2) actions = actions.unsqueeze(-2)
actions = actions.expand(*actions.shape[:-2], states.shape[-2], actions.shape[-1]) actions = actions.expand(*actions.shape[:-2], states.shape[-2], actions.shape[-1])