Improve lot of types, imports and names

This commit is contained in:
henryruhs
2025-03-11 14:43:09 +01:00
parent e33bc0d52a
commit b6b4f9f65b
10 changed files with 64 additions and 64 deletions
@@ -1,5 +1,5 @@
import torch
from torch import Tensor, nn as nn
from torch import Tensor, nn
from face_swapper.src.types import Embedding, TargetAttributes
@@ -19,13 +19,13 @@ class AADGenerator(nn.Module):
def forward(self, target_attributes : TargetAttributes, source_embedding : Embedding) -> Tensor:
feature_map = self.upsample(source_embedding)
feature_map_1 = torch.nn.functional.interpolate(self.res_block_1(feature_map, target_attributes[0], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_2 = torch.nn.functional.interpolate(self.res_block_2(feature_map_1, target_attributes[1], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_3 = torch.nn.functional.interpolate(self.res_block_3(feature_map_2, target_attributes[2], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_4 = torch.nn.functional.interpolate(self.res_block_4(feature_map_3, target_attributes[3], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_5 = torch.nn.functional.interpolate(self.res_block_5(feature_map_4, target_attributes[4], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_6 = torch.nn.functional.interpolate(self.res_block_6(feature_map_5, target_attributes[5], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_7 = torch.nn.functional.interpolate(self.res_block_7(feature_map_6, target_attributes[6], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_1 = nn.functional.interpolate(self.res_block_1(feature_map, target_attributes[0], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_2 = nn.functional.interpolate(self.res_block_2(feature_map_1, target_attributes[1], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_3 = nn.functional.interpolate(self.res_block_3(feature_map_2, target_attributes[2], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_4 = nn.functional.interpolate(self.res_block_4(feature_map_3, target_attributes[3], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_5 = nn.functional.interpolate(self.res_block_5(feature_map_4, target_attributes[4], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_6 = nn.functional.interpolate(self.res_block_6(feature_map_5, target_attributes[5], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
feature_map_7 = nn.functional.interpolate(self.res_block_7(feature_map_6, target_attributes[6], source_embedding), scale_factor = 2, mode = 'bilinear', align_corners = False)
output = self.res_block_8(feature_map_7, target_attributes[7], source_embedding)
return torch.tanh(output)
@@ -41,7 +41,7 @@ class AADLayer(nn.Module):
self.instance_norm = nn.InstanceNorm2d(input_channels)
self.conv_mask = nn.Conv2d(input_channels, 1, kernel_size = 1)
def forward(self, feature_map : Tensor, attribute_embedding : Tensor, id_embedding : Embedding) -> Tensor:
def forward(self, feature_map : Tensor, attribute_embedding : Embedding, id_embedding : Embedding) -> Tensor:
feature_map = self.instance_norm(feature_map)
gamma_attribute = self.conv_gamma(attribute_embedding)
beta_attribute = self.conv_beta(attribute_embedding)
@@ -59,7 +59,7 @@ class AADSequential(nn.Module):
super(AADSequential, self).__init__()
self.layers = nn.ModuleList(args)
def forward(self, feature_map: Tensor, attribute_embedding: Tensor, id_embedding: Embedding) -> Tensor:
def forward(self, feature_map : Tensor, attribute_embedding : Embedding, id_embedding : Embedding) -> Tensor:
for layer in self.layers:
if isinstance(layer, AADLayer):
feature_map = layer(feature_map, attribute_embedding, id_embedding)
@@ -99,7 +99,7 @@ class AADResBlock(nn.Module):
)
self.auxiliary_add_blocks = auxiliary_add_blocks
def forward(self, feature_map : Tensor, attribute_embedding : Tensor, id_embedding : Embedding) -> Tensor:
def forward(self, feature_map : Tensor, attribute_embedding : Embedding, id_embedding : Embedding) -> Tensor:
primary_feature = self.primary_add_blocks(feature_map, attribute_embedding, id_embedding)
if self.input_channels > self.output_channels:
@@ -115,7 +115,7 @@ class PixelShuffleUpsample(nn.Module):
self.conv = nn.Conv2d(in_channels = input_channels, out_channels = output_channels, kernel_size = 3, padding = 1)
self.pixel_shuffle = nn.PixelShuffle(upscale_factor = 2)
def forward(self, temp : Tensor) -> Tensor:
temp = self.conv(temp.view(temp.shape[0], -1, 1, 1))
temp = self.pixel_shuffle(temp)
return temp
def forward(self, input_tensor : Tensor) -> Tensor:
temp_tensor = self.conv(input_tensor.view(input_tensor.shape[0], -1, 1, 1))
temp_tensor = self.pixel_shuffle(temp_tensor)
return temp_tensor
+12 -12
View File
@@ -1,5 +1,5 @@
import torch
from torch import Tensor, nn as nn
from torch import Tensor, nn
from face_swapper.src.types import TargetAttributes, VisionTensor
@@ -11,11 +11,11 @@ class Upsample(nn.Module):
self.batch_norm = nn.BatchNorm2d(output_channels)
self.leaky_relu = nn.LeakyReLU(0.1, inplace = True)
def forward(self, temp : Tensor, skip_tensor : Tensor) -> Tensor:
temp = self.deconv(temp)
temp = self.batch_norm(temp)
temp = self.leaky_relu(temp)
return torch.cat((temp, skip_tensor), dim = 1)
def forward(self, input_tensor : Tensor, skip_tensor : Tensor) -> Tensor:
temp_tensor = self.deconv(input_tensor)
temp_tensor = self.batch_norm(temp_tensor)
temp_tensor = self.leaky_relu(temp_tensor)
return torch.cat((temp_tensor, skip_tensor), dim = 1)
class DownSample(nn.Module):
@@ -25,11 +25,11 @@ class DownSample(nn.Module):
self.batch_norm = nn.BatchNorm2d(output_channels)
self.leaky_relu = nn.LeakyReLU(0.1, inplace = True)
def forward(self, temp : Tensor) -> Tensor:
temp = self.conv(temp)
temp = self.batch_norm(temp)
temp = self.leaky_relu(temp)
return temp
def forward(self, input_tensor : Tensor) -> Tensor:
temp_tensor = self.conv(input_tensor)
temp_tensor = self.batch_norm(temp_tensor)
temp_tensor = self.leaky_relu(temp_tensor)
return temp_tensor
class UNet(nn.Module):
@@ -63,5 +63,5 @@ class UNet(nn.Module):
upsample_feature_4 = self.upsampler_4(upsample_feature_3, downsample_feature_3)
upsample_feature_5 = self.upsampler_5(upsample_feature_4, downsample_feature_2)
upsample_feature_6 = self.upsampler_6(upsample_feature_5, downsample_feature_1)
output = torch.nn.functional.interpolate(upsample_feature_6, scale_factor = 2, mode = 'bilinear', align_corners = False)
output = nn.functional.interpolate(upsample_feature_6, scale_factor = 2, mode = 'bilinear', align_corners = False)
return bottleneck_output, upsample_feature_1, upsample_feature_2, upsample_feature_3, upsample_feature_4, upsample_feature_5, upsample_feature_6, output