MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / configure_optimizers

Method configure_optimizers

sat/sgm/models/autoencoder.py:360–381  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

358 return groups, num_params
359
360 def configure_optimizers(self) -> List[torch.optim.Optimizer]:
361 if self.trainable_ae_params is None:
362 ae_params = self.get_autoencoder_params()
363 else:
364 ae_params, num_ae_params = self.get_param_groups(self.trainable_ae_params, self.ae_optimizer_args)
365 logpy.info(f"Number of trainable autoencoder parameters: {num_ae_params:,}")
366 if self.trainable_disc_params is None:
367 disc_params = self.get_discriminator_params()
368 else:
369 disc_params, num_disc_params = self.get_param_groups(self.trainable_disc_params, self.disc_optimizer_args)
370 logpy.info(f"Number of trainable discriminator parameters: {num_disc_params:,}")
371 opt_ae = self.instantiate_optimizer_from_config(
372 ae_params,
373 default(self.lr_g_factor, 1.0) * self.learning_rate,
374 self.optimizer_config,
375 )
376 opts = [opt_ae]
377 if len(disc_params) > 0:
378 opt_disc = self.instantiate_optimizer_from_config(disc_params, self.learning_rate, self.optimizer_config)
379 opts.append(opt_disc)
380
381 return opts
382
383 @torch.no_grad()
384 def log_images(self, batch: dict, additional_log_kwargs: Optional[Dict] = None, **kwargs) -> dict:

Callers

nothing calls this directly

Calls 6

get_param_groupsMethod · 0.95
appendMethod · 0.80
defaultFunction · 0.50

Tested by

no test coverage detected