Amordia/Interactive-Automatic-Image-Labeling-Platform-Development
0
1from typing import Optional, Dict, Any, List2import torch3import torch.nn as nn4 5# -----------------------------------------------------------------------------6# Blocks7# -----------------------------------------------------------------------------8 9class Conv2d(nn.Module):10 """ Perform a 2D convolution11 12 inputs are [b, c, h, w] where 13 b is the batch size14 c is the number of channels 15 h is the height16 w is the width17 """18 def __init__(self, 19 in_channels: int, 20 out_channels: int, 21 kernel_size: int, 22 padding: int,23 do_activation: bool = True, 24 ):25 super(Conv2d, self).__init__()26 27 conv = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, padding=padding)28 lst = [conv]29 30 if do_activation:31 lst.append(nn.PReLU())32 33 self.conv = nn.Sequential(*lst)34 35 def forward(self, x):36 # x is [B, C, H, W]37 return self.conv(x)38 39# -----------------------------------------------------------------------------40# Network41# -----------------------------------------------------------------------------42 43class _UNet(nn.Module):44 def __init__(self,45 in_channels: int = 1,46 out_channels: int = 1,47 features: List[int] = [64, 64, 64, 64, 64],48 conv_kernel_size: int = 3,49 conv: Optional[nn.Module] = None,50 conv_kwargs: Dict[str,Any] = {}51 ):52 """53 UNet (but can switch out the Conv)54 """55 super(_UNet, self).__init__()56 57 self.in_channels = in_channels58 59 padding = (conv_kernel_size - 1) // 260 61 self.ups = nn.ModuleList()62 self.downs = nn.ModuleList()63 self.pool = nn.MaxPool2d(kernel_size=2, stride=2)64 65 # Down part of U-Net66 for feat in features:67 self.downs.append(68 conv(69 in_channels, feat, kernel_size=conv_kernel_size, padding=padding, **conv_kwargs70 )71 )72 in_channels = feat73 74 # Up part of U-Net75 for feat in reversed(features):76 self.ups.append(nn.UpsamplingBilinear2d(scale_factor=2))77 self.ups.append(78 conv(79 # Factor of 2 is for the skip connections80 feat * 2, feat, kernel_size=conv_kernel_size, padding=padding, **conv_kwargs81 )82 )83 84 self.bottleneck = conv(85 features[-1], features[-1], kernel_size=conv_kernel_size, padding=padding, **conv_kwargs86 )87 self.final_conv = conv(88 features[0], out_channels, kernel_size=1, padding=0, do_activation=False, **conv_kwargs89 )90 91 def forward(self, x: torch.Tensor) -> torch.Tensor:92 skip_connections = []93 for down in self.downs:94 x = down(x)95 skip_connections.append(x)96 x = self.pool(x)97 98 x = self.bottleneck(x)99 skip_connections = skip_connections[::-1]100 101 for idx in range(0, len(self.ups), 2):102 x = self.ups[idx](x)103 skip_connection = skip_connections[idx // 2]104 105 concat_skip = torch.cat((skip_connection, x), dim=1)106 x = self.ups[idx + 1](concat_skip)107 108 return self.final_conv(x)109 110 111class UNet(_UNet):112 """113 Unet with normal conv blocks114 115 input shape: B x C x H x W116 output shape: B x C x H x W 117 """118 def __init__(self, **kwargs) -> None:119 super().__init__(conv=Conv2d, **kwargs)120 121 def forward(self, x: torch.Tensor) -> torch.Tensor:122 return super().forward(x)123 