MCPcopy Create free account
hub / github.com/JunlinHan/DCLGAN / backward_D_basic

Method backward_D_basic

models/dcl_model.py:171–190  ·  view source on GitHub ↗

Calculate GAN loss for the discriminator Parameters: netD (network) -- the discriminator D real (tensor array) -- real images fake (tensor array) -- images generated by a generator Return the discriminator loss. We also call loss_D.ba

(self, netD, real, fake)

Source from the content-addressed store, hash-verified

169 self.idt_B = self.netG_B(self.real_A)
170
171 def backward_D_basic(self, netD, real, fake):
172 """Calculate GAN loss for the discriminator
173 Parameters:
174 netD (network) -- the discriminator D
175 real (tensor array) -- real images
176 fake (tensor array) -- images generated by a generator
177
178 Return the discriminator loss.
179 We also call loss_D.backward() to calculate the gradients.
180 """
181 # Real
182 pred_real = netD(real)
183 loss_D_real = self.criterionGAN(pred_real, True)
184 # Fake
185 pred_fake = netD(fake.detach())
186 loss_D_fake = self.criterionGAN(pred_fake, False)
187 # Combined loss and calculate gradients
188 loss_D = (loss_D_real + loss_D_fake) * 0.5
189 loss_D.backward()
190 return loss_D
191
192 def backward_D_A(self):
193 """Calculate GAN loss for discriminator D_A"""

Callers 2

backward_D_AMethod · 0.95
backward_D_BMethod · 0.95

Calls 1

backwardMethod · 0.80

Tested by

no test coverage detected