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 -1
View File
@@ -2,7 +2,7 @@ import configparser
from typing import List
import numpy
import torch.nn as nn
from torch import nn
from face_swapper.src.types import VisionTensor
+1 -1
View File
@@ -1,7 +1,7 @@
import configparser
from typing import Tuple
import torch.nn as nn
from torch import nn
from face_swapper.src.networks.attribute_modulator import AADGenerator
from face_swapper.src.networks.encoder import UNet
+2 -2
View File
@@ -18,7 +18,7 @@ class FaceSwapperLoss:
landmarker_path = CONFIG.get('training.model', 'landmarker_path')
motion_extractor_path = CONFIG.get('training.model', 'motion_extractor_path')
self.batch_size = CONFIG.getint('training.loader', 'batch_size')
self.mse_loss = torch.nn.MSELoss()
self.mse_loss = nn.MSELoss()
self.id_embedder = torch.jit.load(id_embedder_path, map_location = 'cpu') # type:ignore[no-untyped-call]
self.landmarker = torch.jit.load(landmarker_path, map_location = 'cpu') # type:ignore[no-untyped-call]
self.motion_extractor = torch.jit.load(motion_extractor_path, map_location = 'cpu') # type:ignore[no-untyped-call]
@@ -127,7 +127,7 @@ class FaceSwapperLoss:
def get_face_landmarks(self, vision_tensor : VisionTensor) -> FaceLandmark203:
vision_tensor_norm = (vision_tensor + 1) * 0.5
vision_tensor_norm = torch.nn.functional.interpolate(vision_tensor_norm, size = (224, 224), mode = 'bilinear')
vision_tensor_norm = nn.functional.interpolate(vision_tensor_norm, size = (224, 224), mode = 'bilinear')
landmarks = self.landmarker(vision_tensor_norm)[2].view(-1, 203, 2)
return landmarks