(self, train_generator=True, keep_data_for_visuals=False)
| 196 | return imgs, data |
| 197 | |
| 198 | def optimize_parameters(self, train_generator=True, keep_data_for_visuals=False): |
| 199 | if self._is_train: |
| 200 | # convert tensor to variables |
| 201 | self._B = self._input_real_img.size(0) |
| 202 | self._real_img = Variable(self._input_real_img) |
| 203 | self._real_cond = Variable(self._input_real_cond) |
| 204 | self._desired_cond = Variable(self._input_desired_cond) |
| 205 | |
| 206 | # train D |
| 207 | loss_D, fake_imgs_masked = self._forward_D() |
| 208 | self._optimizer_D.zero_grad() |
| 209 | loss_D.backward() |
| 210 | self._optimizer_D.step() |
| 211 | |
| 212 | loss_D_gp= self._gradinet_penalty_D(fake_imgs_masked) |
| 213 | self._optimizer_D.zero_grad() |
| 214 | loss_D_gp.backward() |
| 215 | self._optimizer_D.step() |
| 216 | |
| 217 | # train G |
| 218 | if train_generator: |
| 219 | loss_G = self._forward_G(keep_data_for_visuals) |
| 220 | self._optimizer_G.zero_grad() |
| 221 | loss_G.backward() |
| 222 | self._optimizer_G.step() |
| 223 | |
| 224 | def _forward_G(self, keep_data_for_visuals): |
| 225 | # generate fake images |
nothing calls this directly
no test coverage detected