Move mask factor to trainer, Refactor helper

This commit is contained in:
henryruhs
2025-06-20 15:15:02 +02:00
parent 338f49c3dc
commit a86497177d
4 changed files with 22 additions and 22 deletions
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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 =
+16 -16
View File
@@ -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
+4 -4
View File
@@ -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