| """ |
| 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_semantic_model() -> LinknetModel: |
| """Factory: LinkNet with 3-channel input and 12 semantic output channels.""" |
| return LinknetModel(Cin=3, N=12) |
|
|