Spaces:
Sleeping
Sleeping
| """ | |
| Author: Juan Pablo Triana Martinez | |
| LinkNet architecture — standalone for HuggingFace Spaces deployment. | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| class LinknetStem(nn.Module): | |
| def __init__(self, m: int = 3, n: int = 64) -> None: | |
| super().__init__() | |
| self.linknet_stem = nn.Sequential( | |
| nn.Conv2d(m, n, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| nn.MaxPool2d(kernel_size=(3, 3), stride=(2, 2), padding=(1, 1)), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.linknet_stem(x) | |
| class LinknetEncoderBlock(nn.Module): | |
| def __init__(self, m: int, n: int) -> None: | |
| super().__init__() | |
| self.convs_blocks_1 = nn.Sequential( | |
| nn.Conv2d(m, n, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| nn.Conv2d(n, n, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| ) | |
| self.skip_conn = nn.Sequential( | |
| nn.Conv2d(m, n, kernel_size=(1, 1), stride=(2, 2), padding=(0, 0), bias=False), | |
| nn.BatchNorm2d(n), | |
| ) | |
| self.convs_block_2 = nn.Sequential( | |
| nn.Conv2d(n, n, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| nn.Conv2d(n, n, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x1 = self.convs_blocks_1(x) | |
| x2 = x1 + self.skip_conn(x) | |
| x3 = self.convs_block_2(x2) | |
| return x3 + x2 | |
| class LinknetDecoderBlock(nn.Module): | |
| def __init__(self, m: int, n: int) -> None: | |
| super().__init__() | |
| self.conv_block_1 = nn.Sequential( | |
| nn.Conv2d(m, m // 4, kernel_size=(1, 1), bias=False), | |
| nn.BatchNorm2d(m // 4), | |
| nn.ReLU(), | |
| ) | |
| self.upsample_block = nn.Sequential( | |
| nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True), | |
| nn.Conv2d(m // 4, m // 4, kernel_size=(3, 3), padding=(1, 1), bias=False), | |
| nn.BatchNorm2d(m // 4), | |
| nn.ReLU(), | |
| ) | |
| self.conv_block_2 = nn.Sequential( | |
| nn.Conv2d(m // 4, n, kernel_size=(1, 1), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.conv_block_1(x) | |
| x = self.upsample_block(x) | |
| return self.conv_block_2(x) | |
| class LinknetReconstructer(nn.Module): | |
| def __init__(self, N: int = 1, m: int = 64, n: int = 32) -> None: | |
| super().__init__() | |
| self.upsample_block_1 = nn.Sequential( | |
| nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True), | |
| nn.Conv2d(m, n, kernel_size=(3, 3), padding=(1, 1), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| ) | |
| self.conv_block = nn.Sequential( | |
| nn.Conv2d(n, n, kernel_size=(3, 3), padding=(1, 1), bias=False), | |
| nn.BatchNorm2d(n), | |
| nn.ReLU(), | |
| ) | |
| self.upsample_block_2 = nn.Sequential( | |
| nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True), | |
| nn.Conv2d(n, N, kernel_size=(3, 3), padding=(1, 1), bias=False), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.upsample_block_1(x) | |
| x = self.conv_block(x) | |
| return self.upsample_block_2(x) | |
| class LinknetModel(nn.Module): | |
| def __init__(self, Cin: int = 3, N: int = 1) -> None: | |
| super().__init__() | |
| self.stem = LinknetStem(m=Cin, n=64) | |
| self.encoder_block_1 = LinknetEncoderBlock(64, 64) | |
| self.encoder_block_2 = LinknetEncoderBlock(64, 128) | |
| self.encoder_block_3 = LinknetEncoderBlock(128, 256) | |
| self.encoder_block_4 = LinknetEncoderBlock(256, 512) | |
| self.decoder_block_4 = LinknetDecoderBlock(512, 256) | |
| self.decoder_block_3 = LinknetDecoderBlock(256, 128) | |
| self.decoder_block_2 = LinknetDecoderBlock(128, 64) | |
| self.decoder_block_1 = LinknetDecoderBlock(64, 64) | |
| self.reconstructer = LinknetReconstructer(N=N, m=64, n=32) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.stem(x) | |
| x1 = self.encoder_block_1(x) | |
| x2 = self.encoder_block_2(x1) | |
| x3 = self.encoder_block_3(x2) | |
| x4 = self.encoder_block_4(x3) | |
| x = self.decoder_block_4(x4) + x3 | |
| x = self.decoder_block_3(x) + x2 | |
| x = self.decoder_block_2(x) + x1 | |
| x = self.decoder_block_1(x) | |
| return self.reconstructer(x) | |
| def create_binary_model() -> LinknetModel: | |
| """Factory: LinkNet with 3-channel input and 1 binary output channel.""" | |
| return LinknetModel(Cin=3, N=1) | |