| 59 | |
| 60 | class SirenNet(nn.Module): |
| 61 | def __init__(self, dim_in, dim_hidden, dim_out, num_layers, w0 = 1., w0_initial = 30., use_bias = True, final_activation = None, degreeinput = False, dropout = False): |
| 62 | super().__init__() |
| 63 | self.num_layers = num_layers |
| 64 | self.dim_hidden = dim_hidden |
| 65 | self.degreeinput = degreeinput |
| 66 | |
| 67 | self.layers = nn.ModuleList([]) |
| 68 | for ind in range(num_layers): |
| 69 | is_first = ind == 0 |
| 70 | layer_w0 = w0_initial if is_first else w0 |
| 71 | layer_dim_in = dim_in if is_first else dim_hidden |
| 72 | |
| 73 | self.layers.append(Siren( |
| 74 | dim_in = layer_dim_in, |
| 75 | dim_out = dim_hidden, |
| 76 | w0 = layer_w0, |
| 77 | use_bias = use_bias, |
| 78 | is_first = is_first, |
| 79 | dropout = dropout |
| 80 | )) |
| 81 | |
| 82 | final_activation = nn.Identity() if not exists(final_activation) else final_activation |
| 83 | self.last_layer = Siren(dim_in = dim_hidden, dim_out = dim_out, w0 = w0, use_bias = use_bias, activation = final_activation, dropout = False) |
| 84 | |
| 85 | def forward(self, x, mods = None): |
| 86 | |