Fix deprecation warning in sigmoid

This commit is contained in:
Johannes Fischer
2021-07-20 12:52:02 +02:00
parent a9c6857b5b
commit ee7b6b607f
2 changed files with 8 additions and 6 deletions

View File

@@ -2,13 +2,11 @@ import torch
from torch.nn import functional from torch.nn import functional
def parse_functional(functional_config): def parse_functional(functional_config):
if functional_config is None: if isinstance(functional_config, str):
return None
elif isinstance(functional_config, str):
if functional_config == 'relu': if functional_config == 'relu':
return functional.relu return functional.relu
elif functional_config == 'sigmoid': elif functional_config == 'sigmoid':
return functional.sigmoid return torch.sigmoid
elif functional_config == 'softmax': elif functional_config == 'softmax':
return functional.softmax return functional.softmax
return None

View File

@@ -24,7 +24,11 @@ def test_constructor():
phi_config["output_dim"] = 2 phi_config["output_dim"] = 2
phi_config["final_activation"] = "sigmoid" phi_config["final_activation"] = "sigmoid"
phi = ds.Phi.from_config(phi_config) 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(): def test_phi():
input_dim = 5 input_dim = 5