(self, discriminator_loss, generator_loss,
hr_discriminator_loss, hr_generator_loss)
| 231 | return trainer |
| 232 | |
| 233 | def prepare_trainer(self, discriminator_loss, generator_loss, |
| 234 | hr_discriminator_loss, hr_generator_loss): |
| 235 | ft_lr_retio = cfg.TRAIN.FT_LR_RETIO |
| 236 | self.discriminator_trainer =\ |
| 237 | self.define_one_trainer(discriminator_loss, |
| 238 | self.discriminator_lr * ft_lr_retio, |
| 239 | 'd_') |
| 240 | self.generator_trainer =\ |
| 241 | self.define_one_trainer(generator_loss, |
| 242 | self.generator_lr * ft_lr_retio, |
| 243 | 'g_') |
| 244 | self.hr_discriminator_trainer =\ |
| 245 | self.define_one_trainer(hr_discriminator_loss, |
| 246 | self.discriminator_lr, |
| 247 | 'hr_d_') |
| 248 | self.hr_generator_trainer =\ |
| 249 | self.define_one_trainer(hr_generator_loss, |
| 250 | self.generator_lr, |
| 251 | 'hr_g_') |
| 252 | |
| 253 | self.ft_generator_trainer = \ |
| 254 | self.define_one_trainer(hr_generator_loss, |
| 255 | self.generator_lr * cfg.TRAIN.FT_LR_RETIO, |
| 256 | 'g_') |
| 257 | |
| 258 | self.log_vars.append(("hr_d_learning_rate", self.discriminator_lr)) |
| 259 | self.log_vars.append(("hr_g_learning_rate", self.generator_lr)) |
| 260 | |
| 261 | def define_summaries(self): |
| 262 | '''Helper function for init_opt''' |
no test coverage detected