Update deepsets phi module to allow different final activation

This commit is contained in:
Johannes Fischer
2021-07-19 18:35:07 +02:00
parent 2ff43b293b
commit 0e6b1e102e
3 changed files with 34 additions and 4 deletions

View File

@@ -1,6 +1,8 @@
import torch import torch
from torch import nn from torch import nn
from src.nets.util import parse_functional
class DeepSetsModule(nn.Module): 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): 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): 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 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.hidden_layers = [nn.Linear(hidden_dim, hidden_dim) for _ in range(hidden_n - 1)]
# self.out_layer = nn.Linear(hidden_dim, output_dim) # self.out_layer = nn.Linear(hidden_dim, output_dim)
self.activation = nn.functional.relu self.activation = nn.functional.relu
self.final_activation = final_activation if final_activation else self.activation
def forward(self, x): def forward(self, x):
for layer in self.layers: for layer in self.layers[:-1]:
x = self.activation(layer(x)) x = self.activation(layer(x))
x = self.final_activation(self.layers[-1](x))
return x return x
@staticmethod @staticmethod
def from_config(config): 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)

14
src/nets/util.py Normal file
View File

@@ -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

View File

@@ -1,6 +1,7 @@
import torch import torch
import random import random
from src.nets import deepsets as ds from src.nets import deepsets as ds
import copy
ds_config = { ds_config = {
"input_dim": 5, "input_dim": 5,
@@ -13,11 +14,17 @@ ds_config = {
"hidden_n": 1, "hidden_n": 1,
"hidden_dim": 10, "hidden_dim": 10,
}, },
"output_dim" : 1, "output_dim" : 2,
} }
def test_constructor(): def test_constructor():
m = ds.DeepSetsModule.from_config(ds_config) 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(): def test_phi():
input_dim = 5 input_dim = 5