(self)
| 32 | self._init_losses() |
| 33 | |
| 34 | def _init_create_networks(self): |
| 35 | # generator network |
| 36 | self._G = self._create_generator() |
| 37 | self._G.init_weights() |
| 38 | if len(self._gpu_ids) > 1: |
| 39 | self._G = torch.nn.DataParallel(self._G, device_ids=self._gpu_ids) |
| 40 | self._G.cuda() |
| 41 | |
| 42 | # discriminator network |
| 43 | self._D = self._create_discriminator() |
| 44 | self._D.init_weights() |
| 45 | if len(self._gpu_ids) > 1: |
| 46 | self._D = torch.nn.DataParallel(self._D, device_ids=self._gpu_ids) |
| 47 | self._D.cuda() |
| 48 | |
| 49 | def _create_generator(self): |
| 50 | return NetworksFactory.get_by_name('generator_wasserstein_gan', c_dim=self._opt.cond_nc) |
no test coverage detected