| 25 | ]) |
| 26 | |
| 27 | def forward(self, sample_id, class_label=None, original_image=None): |
| 28 | # Eval mode for the source and target gan_wrapper. |
| 29 | self.source_gan_wrapper.eval() |
| 30 | self.target_gan_wrapper.eval() |
| 31 | |
| 32 | assert not self.training |
| 33 | |
| 34 | if getattr(self.source_gan_wrapper, "model_embedding_space", False): |
| 35 | assert getattr(self.target_gan_wrapper, "model_embedding_space", False) |
| 36 | assert not getattr(self.source_gan_wrapper, "enforce_class_input", False) |
| 37 | assert not getattr(self.target_gan_wrapper, "enforce_class_input", False) |
| 38 | raise NotImplementedError() |
| 39 | elif getattr(self.source_gan_wrapper, "enforce_class_input", False): |
| 40 | assert getattr(self.target_gan_wrapper, "enforce_class_input", False) |
| 41 | assert not getattr(self.source_gan_wrapper, "model_embedding_space", False) |
| 42 | assert not getattr(self.target_gan_wrapper, "model_embedding_space", False) |
| 43 | assert class_label is not None |
| 44 | z = self.source_gan_wrapper.encode(image=original_image, class_label=class_label) |
| 45 | img = self.target_gan_wrapper(z=z, class_label=class_label) |
| 46 | else: |
| 47 | assert class_label is None |
| 48 | z = self.source_gan_wrapper.encode(image=original_image) |
| 49 | img = self.target_gan_wrapper(z=z) |
| 50 | |
| 51 | # Placeholders |
| 52 | losses = dict() |
| 53 | weighted_loss = torch.zeros_like(sample_id).float() |
| 54 | |
| 55 | return (original_image, img), weighted_loss, losses |
| 56 | |
| 57 | @property |
| 58 | def device(self): |