diff --git a/src/nets/util.py b/src/nets/util.py index 71a2150..e868adc 100644 --- a/src/nets/util.py +++ b/src/nets/util.py @@ -2,13 +2,11 @@ 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 isinstance(functional_config, str): if functional_config == 'relu': return functional.relu elif functional_config == 'sigmoid': - return functional.sigmoid + return torch.sigmoid elif functional_config == 'softmax': return functional.softmax - \ No newline at end of file + return None \ No newline at end of file diff --git a/tests/nets/test_deepsets.py b/tests/nets/test_deepsets.py index d47d969..2321918 100644 --- a/tests/nets/test_deepsets.py +++ b/tests/nets/test_deepsets.py @@ -24,7 +24,11 @@ def test_constructor(): 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 + assert phi.final_activation == torch.sigmoid + + phi_config["final_activation"] = "relu" + phi = ds.Phi.from_config(phi_config) + assert phi.final_activation == torch.nn.functional.relu def test_phi(): input_dim = 5