Update deepsets phi module to allow different final activation
This commit is contained in:
@@ -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
14
src/nets/util.py
Normal 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
|
||||||
|
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user