Spaces:
Build error
Build error
| """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) | |