From 827a8e717271730eec59e86b4785b5f28b9a7c86 Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Thu, 22 Jul 2021 14:33:05 +0200 Subject: [PATCH] bugfix in Phi module nn.ModuleList has to be used in order to register layer parameters as module parameters (similar to add_module) --- src/nets/deepsets.py | 5 +---- tests/nets/test_deepsets.py | 2 ++ 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/src/nets/deepsets.py b/src/nets/deepsets.py index d1f3bc0..a719eee 100644 --- a/src/nets/deepsets.py +++ b/src/nets/deepsets.py @@ -90,13 +90,10 @@ class Phi(nn.Module): super(Phi, self).__init__() self.input_dim = input_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): self.layers.append(nn.Linear(hidden_dim, hidden_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.final_activation = final_activation if final_activation else lambda x: x diff --git a/tests/nets/test_deepsets.py b/tests/nets/test_deepsets.py index 7726f57..c5eb903 100644 --- a/tests/nets/test_deepsets.py +++ b/tests/nets/test_deepsets.py @@ -41,6 +41,8 @@ def test_phi(): y = phi(torch.rand(input_dim)) y = phi(torch.rand(7,7,7,input_dim)) + assert len(phi.parameters() > 0) + def test_deepsets(): m = ds.DeepSetsModule.from_config(ds_config)