medimaging's picture
Upload 2 files
288bc38 verified
Raw
History Blame
2.94 kB
# Based on https://github.com/lucidrains/siren-pytorch
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)