bugfix in Phi module

nn.ModuleList has to be used in order to register layer parameters as module parameters (similar to add_module)
This commit is contained in:
Johannes Fischer
2021-07-22 14:33:05 +02:00
parent 1ca9914bf9
commit 827a8e7172
2 changed files with 3 additions and 4 deletions

View File

@@ -90,13 +90,10 @@ class Phi(nn.Module):
super(Phi, self).__init__() super(Phi, self).__init__()
self.input_dim = input_dim self.input_dim = input_dim
self.output_dim = output_dim self.output_dim = output_dim
self.layers = [nn.Linear(self.input_dim, hidden_dim)] self.layers = nn.ModuleList([nn.Linear(self.input_dim, hidden_dim)])
for _ in range(hidden_n - 1): for _ in range(hidden_n - 1):
self.layers.append(nn.Linear(hidden_dim, hidden_dim)) self.layers.append(nn.Linear(hidden_dim, hidden_dim))
self.layers.append(nn.Linear(hidden_dim, self.output_dim)) self.layers.append(nn.Linear(hidden_dim, self.output_dim))
# self.in_layer = nn.Linear(input_dim, hidden_dim)
# 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.activation = nn.functional.relu
self.final_activation = final_activation if final_activation else lambda x: x self.final_activation = final_activation if final_activation else lambda x: x

View File

@@ -41,6 +41,8 @@ def test_phi():
y = phi(torch.rand(input_dim)) y = phi(torch.rand(input_dim))
y = phi(torch.rand(7,7,7,input_dim)) y = phi(torch.rand(7,7,7,input_dim))
assert len(phi.parameters() > 0)
def test_deepsets(): def test_deepsets():
m = ds.DeepSetsModule.from_config(ds_config) m = ds.DeepSetsModule.from_config(ds_config)