add erode for export and make it conditional

This commit is contained in:
harisreedhar
2025-06-19 16:29:30 +05:30
parent bd762c4c38
commit 56b71048e3
3 changed files with 19 additions and 5 deletions
+10 -3
View File
@@ -69,7 +69,14 @@ def resolve_static_file_pattern(file_pattern : str) -> List[str]:
def dilate_mask(input_tensor : Tensor, factor : float) -> Tensor:
padding = round(input_tensor.shape[2] * factor)
kernel_size = 1 + 2 * padding
kernel = torch.ones((1, 1, kernel_size, kernel_size), dtype = input_tensor.dtype, device = input_tensor.device)
dilate_tensor = nn.functional.conv2d(input_tensor, kernel, padding = padding)
dilate_tensor = torch.sigmoid(2 * (dilate_tensor - 0.5))
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 -1
View File
@@ -170,7 +170,10 @@ class MaskLoss(nn.Module):
def forward(self, target_tensor : Tensor, output_mask : Mask) -> Tuple[Loss, Loss]:
target_mask = self.calc_mask(target_tensor)
target_mask = dilate_mask(target_mask, self.config_mask_dilate)
if self.config_mask_dilate > 0:
target_mask = dilate_mask(target_mask, self.config_mask_dilate)
target_mask = target_mask.view(-1, self.config_output_size, self.config_output_size)
output_mask = output_mask.view(-1, self.config_output_size, self.config_output_size)
mask_loss = self.mse_loss(target_mask, output_mask)
+5 -1
View File
@@ -14,7 +14,7 @@ from torch.utils.data import ConcatDataset, Dataset, random_split
from torchdata.stateful_dataloader import StatefulDataLoader
from .dataset import DynamicDataset
from .helper import apply_noise, calc_embedding, overlay_mask
from .helper import apply_noise, calc_embedding, erode_mask, overlay_mask
from .models.discriminator import Discriminator
from .models.generator import Generator
from .models.loss import AdversarialLoss, CycleLoss, DiscriminatorLoss, FeatureLoss, GazeLoss, IdentityLoss, MaskLoss, ReconstructionLoss
@@ -45,6 +45,7 @@ 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()
@@ -66,6 +67,9 @@ 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)
return output_tensor, output_mask
def configure_optimizers(self) -> Tuple[OptimizerSet, OptimizerSet]: