MCPcopy Create free account
hub / github.com/ChenWu98/cycle-diffusion / forward

Method forward

model/unsupervised_translation.py:27–55  ·  view source on GitHub ↗
(self, sample_id, class_label=None, original_image=None)

Source from the content-addressed store, hash-verified

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):

Callers

nothing calls this directly

Calls 1

encodeMethod · 0.45

Tested by

no test coverage detected