File size: 4,999 Bytes
82ddd6b | 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 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | """
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)
|