"""TransformNet from ebylmz/fast-neural-style-transfer (MIT).""" import torch import torch.nn as nn class ConvLayer(nn.Module): def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int, relu: bool = True): super().__init__() layers = [ nn.Conv2d( in_channels, out_channels, kernel_size, stride, padding=kernel_size // 2, padding_mode="reflect", ), nn.InstanceNorm2d(out_channels, affine=True), ] if relu: layers.append(nn.ReLU(inplace=True)) self.block = nn.Sequential(*layers) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.block(x) class ResidualBlock(nn.Module): def __init__(self, channels: int): super().__init__() self.block = nn.Sequential( ConvLayer(channels, channels, kernel_size=3, stride=1, relu=True), ConvLayer(channels, channels, kernel_size=3, stride=1, relu=False), ) def forward(self, x: torch.Tensor) -> torch.Tensor: return x + self.block(x) class UpsampleConvLayer(nn.Module): def __init__( self, in_channels: int, out_channels: int, kernel_size: int, stride: int = 1, upsample: int | None = None, ): super().__init__() layers: list[nn.Module] = [] if upsample: layers.append(nn.Upsample(scale_factor=upsample, mode="nearest")) layers.extend( [ nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding=kernel_size // 2), nn.InstanceNorm2d(out_channels, affine=True), nn.ReLU(inplace=True), ] ) self.block = nn.Sequential(*layers) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.block(x) class TransformNet(nn.Module): def __init__(self): super().__init__() self.downsampling = nn.Sequential( ConvLayer(3, 32, kernel_size=9, stride=1), ConvLayer(32, 64, kernel_size=3, stride=2), ConvLayer(64, 128, kernel_size=3, stride=2), ) self.residuals = nn.Sequential( ResidualBlock(128), ResidualBlock(128), ResidualBlock(128), ResidualBlock(128), ResidualBlock(128), ) self.upsampling = nn.Sequential( UpsampleConvLayer(128, 64, kernel_size=3, upsample=2), UpsampleConvLayer(64, 32, kernel_size=3, upsample=2), nn.Conv2d(32, 3, kernel_size=9, stride=1, padding=4, padding_mode="reflect"), ) def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.downsampling(x) x = self.residuals(x) return self.upsampling(x)