(self)
| 373 | return groups, num_params |
| 374 | |
| 375 | def configure_optimizers(self) -> List[torch.optim.Optimizer]: |
| 376 | if self.trainable_ae_params is None: |
| 377 | ae_params = self.get_autoencoder_params() |
| 378 | else: |
| 379 | ae_params, num_ae_params = self.get_param_groups(self.trainable_ae_params, self.ae_optimizer_args) |
| 380 | logpy.info(f"Number of trainable autoencoder parameters: {num_ae_params:,}") |
| 381 | if self.trainable_disc_params is None: |
| 382 | disc_params = self.get_discriminator_params() |
| 383 | else: |
| 384 | disc_params, num_disc_params = self.get_param_groups(self.trainable_disc_params, self.disc_optimizer_args) |
| 385 | logpy.info(f"Number of trainable discriminator parameters: {num_disc_params:,}") |
| 386 | opt_ae = self.instantiate_optimizer_from_config( |
| 387 | ae_params, |
| 388 | default(self.lr_g_factor, 1.0) * self.learning_rate, |
| 389 | self.optimizer_config, |
| 390 | ) |
| 391 | opts = [opt_ae] |
| 392 | if len(disc_params) > 0: |
| 393 | opt_disc = self.instantiate_optimizer_from_config(disc_params, self.learning_rate, self.optimizer_config) |
| 394 | opts.append(opt_disc) |
| 395 | |
| 396 | return opts |
| 397 | |
| 398 | @torch.no_grad() |
| 399 | def log_images(self, batch: dict, additional_log_kwargs: Optional[Dict] = None, **kwargs) -> dict: |
nothing calls this directly
no test coverage detected