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

Method configure_optimizers

sat/vae_modules/autoencoder.py:375–396  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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:

Callers

nothing calls this directly

Calls 6

get_param_groupsMethod · 0.95
defaultFunction · 0.90
appendMethod · 0.80

Tested by

no test coverage detected