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)