diff --git a/src/nets/deepsets.py b/src/nets/deepsets.py index 35afc3b..edbb4a0 100644 --- a/src/nets/deepsets.py +++ b/src/nets/deepsets.py @@ -1,6 +1,8 @@ import torch from torch import nn +from src.nets.util import parse_functional + class DeepSetsModule(nn.Module): def __init__(self, input_dim, phi_hidden_n, phi_hidden_dim, latent_dim, rho_hidden_n, rho_hidden_dim, output_dim): """ @@ -70,7 +72,7 @@ class DeepSetsModule(nn.Module): class Phi(nn.Module): - def __init__(self, input_dim, hidden_n, hidden_dim, output_dim): + def __init__(self, input_dim, hidden_n, hidden_dim, output_dim, final_activation=None): """ Fully connected feedforward network with same size for all hidden layers and ReLU activation @@ -91,12 +93,19 @@ class Phi(nn.Module): # self.hidden_layers = [nn.Linear(hidden_dim, hidden_dim) for _ in range(hidden_n - 1)] # self.out_layer = nn.Linear(hidden_dim, output_dim) self.activation = nn.functional.relu + self.final_activation = final_activation if final_activation else self.activation def forward(self, x): - for layer in self.layers: + for layer in self.layers[:-1]: x = self.activation(layer(x)) + x = self.final_activation(self.layers[-1](x)) return x @staticmethod def from_config(config): - return Phi(config["input_dim"], config["hidden_n"], config["hidden_dim"], config["output_dim"]) + args = (config["input_dim"], config["hidden_n"], config["hidden_dim"], config["output_dim"]) + if "final_activation" in config: + kwargs = {"final_activation": parse_functional(config["final_activation"])} + else: + kwargs = {} + return Phi(*args, **kwargs) diff --git a/src/nets/util.py b/src/nets/util.py new file mode 100644 index 0000000..71a2150 --- /dev/null +++ b/src/nets/util.py @@ -0,0 +1,14 @@ +import torch +from torch.nn import functional + +def parse_functional(functional_config): + if functional_config is None: + return None + elif isinstance(functional_config, str): + if functional_config == 'relu': + return functional.relu + elif functional_config == 'sigmoid': + return functional.sigmoid + elif functional_config == 'softmax': + return functional.softmax + \ No newline at end of file diff --git a/tests/nets/test_deepsets.py b/tests/nets/test_deepsets.py index 2086044..d47d969 100644 --- a/tests/nets/test_deepsets.py +++ b/tests/nets/test_deepsets.py @@ -1,6 +1,7 @@ import torch import random from src.nets import deepsets as ds +import copy ds_config = { "input_dim": 5, @@ -13,11 +14,17 @@ ds_config = { "hidden_n": 1, "hidden_dim": 10, }, - "output_dim" : 1, + "output_dim" : 2, } def test_constructor(): m = ds.DeepSetsModule.from_config(ds_config) + phi_config = copy.deepcopy(ds_config["phi"]) + phi_config["input_dim"] = 5 + phi_config["output_dim"] = 2 + phi_config["final_activation"] = "sigmoid" + phi = ds.Phi.from_config(phi_config) + assert phi.final_activation == torch.nn.functional.sigmoid def test_phi(): input_dim = 5