Train MaskNet based on the output

This commit is contained in:
henryruhs
2025-03-12 21:35:07 +01:00
parent 0732924f1e
commit cf0bd93814
3 changed files with 26 additions and 32 deletions
+2 -2
View File
@@ -34,8 +34,8 @@ class MaskNet(nn.Module):
UpSample(num_filters, num_filters)
])
def forward(self, target_tensor : Tensor, target_attribute : Attribute) -> Tensor:
output_tensor = torch.cat([ target_tensor, target_attribute ], dim = 1)
def forward(self, input_tensor : Tensor, input_attribute : Attribute) -> Tensor:
output_tensor = torch.cat([ input_tensor, input_attribute ], dim = 1)
for down_sample in self.down_samples:
output_tensor = down_sample(output_tensor)