| |
| import torch |
| from torch import nn |
| from math import sqrt |
|
|
|
|
| class Sine(nn.Module): |
| """Sine activation with scaling.""" |
| def __init__(self, w0=1.): |
| super().__init__() |
| self.w0 = w0 |
|
|
| def forward(self, x): |
| return torch.sin(self.w0 * x) |
|
|
|
|
| class SirenLayer(nn.Module): |
| def __init__(self, dim_in, dim_out, w0=30., c=6., is_first=False, |
| use_bias=True, activation=None): |
| super().__init__() |
| self.dim_in = dim_in |
| self.is_first = is_first |
| self.linear = nn.Linear(dim_in, dim_out, bias=use_bias) |
| w_std = (1 / dim_in) if self.is_first else (sqrt(c / dim_in) / w0) |
| nn.init.uniform_(self.linear.weight, -w_std, w_std) |
| if use_bias: |
| nn.init.uniform_(self.linear.bias, -w_std, w_std) |
| self.activation = Sine(w0) if activation is None else activation |
|
|
| def forward(self, x): |
| return self.activation(self.linear(x)) |
|
|
|
|
| class Siren(nn.Module): |
| def __init__(self, dim_in, dim_hidden, dim_out, num_layers, w0=30., |
| w0_initial=30., use_bias=True, final_activation=None): |
| super().__init__() |
| layers = [] |
| for ind in range(num_layers): |
| is_first = ind == 0 |
| layer_w0 = w0_initial if is_first else w0 |
| layer_dim_in = dim_in if is_first else dim_hidden |
| layers.append(SirenLayer(dim_in=layer_dim_in, dim_out=dim_hidden, |
| w0=layer_w0, use_bias=use_bias, is_first=is_first)) |
| self.net = nn.Sequential(*layers) |
| final_activation = nn.Identity() if final_activation is None else final_activation |
| self.last_layer = SirenLayer(dim_in=dim_hidden, dim_out=dim_out, w0=w0, |
| use_bias=use_bias, activation=final_activation) |
|
|
| def forward(self, x): |
| return self.last_layer(self.net(x)) |
|
|
|
|
| class MLP(nn.Module): |
| def __init__(self, dim_in, dim_hidden, dim_out, num_layers, |
| activation=None, siren_start=False, siren_end=False): |
| super().__init__() |
| if activation is None: |
| activation = nn.ReLU() |
| fc = [] |
| if not siren_start: |
| fc.extend([nn.Linear(dim_in, dim_hidden), activation]) |
| else: |
| fc.append(SirenLayer(dim_in=dim_in, dim_out=dim_hidden, |
| w0=30, use_bias=True, is_first=True)) |
| for _ in range(num_layers - 1): |
| fc.extend([nn.Linear(dim_hidden, dim_hidden), activation]) |
| if siren_end: |
| fc.append(SirenLayer(dim_in=dim_hidden, dim_out=dim_hidden, |
| w0=30, use_bias=True, is_first=False)) |
| else: |
| fc.extend([nn.Linear(dim_hidden, dim_hidden), activation]) |
| self.encoder = nn.Sequential(*fc, nn.Linear(dim_hidden, dim_out)) |
|
|
| def forward(self, x): |
| return self.encoder(x) |
|
|