Parametrize discriminator architecture
This commit is contained in:
@@ -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,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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])
|
||||||
|
|||||||
Reference in New Issue
Block a user