From a86497177d95a840df8d195015e78b3e5551b715 Mon Sep 17 00:00:00 2001 From: henryruhs Date: Fri, 20 Jun 2025 15:15:02 +0200 Subject: [PATCH] Move mask factor to trainer, Refactor helper --- hyperswap/README.md | 2 +- hyperswap/config.ini | 2 +- hyperswap/src/helper.py | 32 ++++++++++++++++---------------- hyperswap/src/training.py | 8 ++++---- 4 files changed, 22 insertions(+), 22 deletions(-) diff --git a/hyperswap/README.md b/hyperswap/README.md index 0f68da9..2d0b4e8 100644 --- a/hyperswap/README.md +++ b/hyperswap/README.md @@ -84,13 +84,13 @@ reconstruction_weight = 10.0 identity_weight = 20.0 gaze_weight = 0.05 mask_weight = 5.0 -mask_dilate = 0.01 ``` ``` [training.trainer] accumulate_size = 4 gradient_clip = 20.0 +mask_factor = 0.01 noise_factor = 0.05 max_epochs = 50 strategy = auto diff --git a/hyperswap/config.ini b/hyperswap/config.ini index f06ae24..199094c 100644 --- a/hyperswap/config.ini +++ b/hyperswap/config.ini @@ -44,11 +44,11 @@ reconstruction_weight = identity_weight = gaze_weight = mask_weight = -mask_dilate = [training.trainer] accumulate_size = gradient_clip = +mask_factor = noise_factor = max_epochs = strategy = diff --git a/hyperswap/src/helper.py b/hyperswap/src/helper.py index 3741b4d..71ab19b 100644 --- a/hyperswap/src/helper.py +++ b/hyperswap/src/helper.py @@ -55,6 +55,22 @@ def overlay_mask(input_tensor : Tensor, input_mask : Mask) -> Tensor: return output_tensor +def dilate_mask(input_tensor : Tensor, factor : float) -> Tensor: + padding = round(input_tensor.shape[2] * factor) + kernel_size = 1 + 2 * padding + temp_tensor = nn.functional.pad(input_tensor, (padding, padding, padding, padding), mode = 'replicate') + output_tensor = nn.functional.max_pool2d(temp_tensor, kernel_size = kernel_size, stride = 1, padding = 0) + return output_tensor + + +def erode_mask(input_tensor : Tensor, factor : float) -> Tensor: + padding = round(input_tensor.shape[2] * factor) + kernel_size = 1 + 2 * padding + temp_tensor = 1 - nn.functional.pad(input_tensor, (padding, padding, padding, padding), mode = 'replicate') + output_tensor = 1 - nn.functional.max_pool2d(temp_tensor, kernel_size = kernel_size, stride = 1, padding = 0) + return output_tensor + + def apply_noise(input_tensor : Tensor, factor : float) -> Tensor: noise_tensor = torch.randn_like(input_tensor) * factor output_tensor = input_tensor + noise_tensor @@ -64,19 +80,3 @@ def apply_noise(input_tensor : Tensor, factor : float) -> Tensor: @lru_cache(maxsize = None) def resolve_static_file_pattern(file_pattern : str) -> List[str]: return sorted(glob.glob(file_pattern)) - - -def dilate_mask(input_tensor : Tensor, factor : float) -> Tensor: - padding = round(input_tensor.shape[2] * factor) - kernel_size = 1 + 2 * padding - pad_tensor = nn.functional.pad(input_tensor, (padding, padding, padding, padding), mode = 'replicate') - dilate_tensor = nn.functional.max_pool2d(pad_tensor, kernel_size = kernel_size, stride = 1, padding = 0) - return dilate_tensor - - -def erode_mask(input_tensor : Tensor, factor : float) -> Tensor: - padding = round(input_tensor.shape[2] * factor) - kernel_size = 1 + 2 * padding - pad_tensor = 1 - nn.functional.pad(input_tensor, (padding, padding, padding, padding), mode = 'replicate') - dilate_tensor = 1 - nn.functional.max_pool2d(pad_tensor, kernel_size = kernel_size, stride = 1, padding = 0) - return dilate_tensor diff --git a/hyperswap/src/training.py b/hyperswap/src/training.py index 6aac4fc..d923c9a 100644 --- a/hyperswap/src/training.py +++ b/hyperswap/src/training.py @@ -33,9 +33,10 @@ class HyperSwapTrainer(LightningModule): self.config_loss_embedder_path = config_parser.get('training.model', 'loss_embedder_path') self.config_gazer_path = config_parser.get('training.model', 'gazer_path') self.config_face_masker_path = config_parser.get('training.model', 'face_masker_path') - self.config_noise_factor = config_parser.getfloat('training.trainer', 'noise_factor') self.config_accumulate_size = config_parser.getfloat('training.trainer', 'accumulate_size') self.config_gradient_clip = config_parser.getfloat('training.trainer', 'gradient_clip') + self.config_mask_factor = config_parser.getfloat('training.trainer', 'mask_factor') + self.config_noise_factor = config_parser.getfloat('training.trainer', 'noise_factor') self.config_preview_frequency = config_parser.getint('training.trainer', 'preview_frequency') self.config_generator_learning_rate = config_parser.getfloat('training.optimizer.generator', 'learning_rate') self.config_generator_momentum = config_parser.getfloat('training.optimizer.generator', 'momentum') @@ -45,7 +46,6 @@ class HyperSwapTrainer(LightningModule): self.config_discriminator_momentum = config_parser.getfloat('training.optimizer.discriminator', 'momentum') self.config_discriminator_scheduler_factor = config_parser.getfloat('training.optimizer.discriminator', 'scheduler_factor') self.config_discriminator_scheduler_patience = config_parser.getint('training.optimizer.discriminator', 'scheduler_patience') - self.config_mask_dilate = config_parser.getfloat('training.losses', 'mask_dilate') self.generator_embedder = torch.jit.load(self.config_generator_embedder_path, map_location = 'cpu').eval() self.loss_embedder = torch.jit.load(self.config_loss_embedder_path, map_location = 'cpu').eval() self.gazer = torch.jit.load(self.config_gazer_path, map_location = 'cpu').eval() @@ -67,8 +67,8 @@ class HyperSwapTrainer(LightningModule): generator_target_features = self.generator.encode_features(target_tensor) output_tensor, output_mask = self.generator(source_embedding, target_tensor, generator_target_features) - if self.config_mask_dilate > 0: - output_mask = erode_mask(output_mask, self.config_mask_dilate) + if self.config_mask_factor > 0: + output_mask = erode_mask(output_mask, self.config_mask_factor) return output_tensor, output_mask