Team Ai
Apppublic

Amordia/Interactive-Automatic-Image-Labeling-Platform-Development

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
network.py123 linesDownload Raw Back to root
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