MCPcopy Create free account
hub / github.com/albertpumarola/GANimation / optimize_parameters

Method optimize_parameters

models/ganimation.py:198–222  ·  view source on GitHub ↗
(self, train_generator=True, keep_data_for_visuals=False)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 3

_forward_DMethod · 0.95
_gradinet_penalty_DMethod · 0.95
_forward_GMethod · 0.95

Tested by

no test coverage detected