File size: 1,759 Bytes
8880eca | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 | import torch
import torch.nn as nn
from onescience.modules.decoder.unet_decoder import UNetDecoder2D
from onescience.modules.encoder.unet_encoder import UNetEncoder2D
from onescience.modules.head.unet_head import UNetHead2D
class DecoderPath(nn.Module):
def __init__(self, base_channels, num_stages, bilinear, normtype, kernel_size):
super().__init__()
self.decoder = UNetDecoder2D(
base_channels=base_channels,
num_stages=num_stages,
bilinear=bilinear,
normtype=normtype,
kernel_size=kernel_size,
)
self.head = UNetHead2D(
in_channels=base_channels,
out_channels=1,
kernel_size=1,
)
def forward(self, features):
return self.head(self.decoder(features))
class UNetEx(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
base_channels: int = 16,
num_stages: int = 2,
bilinear: bool = True,
normtype: str = "bn",
kernel_size: int = 3,
):
super().__init__()
self.encoder = UNetEncoder2D(
in_channels=in_channels,
base_channels=base_channels,
num_stages=num_stages,
bilinear=bilinear,
normtype=normtype,
kernel_size=kernel_size,
)
self.decoders = nn.ModuleList(
[
DecoderPath(base_channels, num_stages, bilinear, normtype, kernel_size)
for _ in range(out_channels)
]
)
def forward(self, x):
features = self.encoder(x)
outputs = [decoder_path(features) for decoder_path in self.decoders]
return torch.cat(outputs, dim=1)
|