juanpajedrez
Initial commit: all files for app building PDF semantic segmentation
82ddd6b
Raw
History Blame Contribute Delete
5 kB
"""
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)