This commit is contained in:
harisreedhar
2025-03-11 14:43:10 +01:00
committed by henryruhs
parent 8f1f002c64
commit 1659805b08
5 changed files with 47 additions and 28 deletions
+1 -1
View File
@@ -14,7 +14,7 @@ class Generator(nn.Module):
super().__init__()
self.encoder = UNet(config_parser)
self.generator = AAD(config_parser)
self.masker = MaskNet(67, 1, 16)
self.masker = MaskNet(config_parser)
self.encoder.apply(init_weight)
self.generator.apply(init_weight)
self.masker.apply(init_weight)
+2 -2
View File
@@ -195,12 +195,12 @@ class MaskLoss(nn.Module):
def calc_mask(self, target_tensor : Tensor) -> Tensor:
target_tensor = torch.nn.functional.interpolate(target_tensor, (512, 512), mode = 'bilinear')
face_indices = torch.tensor([ 1, 2, 3, 4, 5, 10, 11, 12, 13 ]).to(target_tensor.device)
face_mask_regions = torch.tensor([ 1, 2, 3, 4, 5, 10, 11, 12, 13 ]).to(target_tensor.device)
with torch.no_grad():
output_tensor = self.parser(target_tensor)[0]
output_tensor = output_tensor.argmax(1)
output_tensor = torch.isin(output_tensor, face_indices).to(target_tensor.dtype)
output_tensor = torch.isin(output_tensor, face_mask_regions).to(target_tensor.dtype)
output_tensor = output_tensor.view(-1, 1, 512, 512)
output_tensor = torch.nn.functional.interpolate(output_tensor, (self.config_output_size, self.config_output_size), mode = 'bilinear')
return output_tensor