| 117 | """ |
| 118 | |
| 119 | def __init__( |
| 120 | self, |
| 121 | *args, |
| 122 | encoder_config: Dict, |
| 123 | decoder_config: Dict, |
| 124 | loss_config: Dict, |
| 125 | regularizer_config: Dict, |
| 126 | optimizer_config: Union[Dict, None] = None, |
| 127 | lr_g_factor: float = 1.0, |
| 128 | trainable_ae_params: Optional[List[List[str]]] = None, |
| 129 | ae_optimizer_args: Optional[List[dict]] = None, |
| 130 | trainable_disc_params: Optional[List[List[str]]] = None, |
| 131 | disc_optimizer_args: Optional[List[dict]] = None, |
| 132 | disc_start_iter: int = 0, |
| 133 | diff_boost_factor: float = 3.0, |
| 134 | ckpt_engine: Union[None, str, dict] = None, |
| 135 | ckpt_path: Optional[str] = None, |
| 136 | additional_decode_keys: Optional[List[str]] = None, |
| 137 | **kwargs, |
| 138 | ): |
| 139 | super().__init__(*args, **kwargs) |
| 140 | self.automatic_optimization = False # pytorch lightning |
| 141 | |
| 142 | self.encoder: torch.nn.Module = instantiate_from_config(encoder_config) |
| 143 | self.decoder: torch.nn.Module = instantiate_from_config(decoder_config) |
| 144 | self.loss: torch.nn.Module = instantiate_from_config(loss_config) |
| 145 | self.regularization: AbstractRegularizer = instantiate_from_config(regularizer_config) |
| 146 | self.optimizer_config = default(optimizer_config, {"target": "torch.optim.Adam"}) |
| 147 | self.diff_boost_factor = diff_boost_factor |
| 148 | self.disc_start_iter = disc_start_iter |
| 149 | self.lr_g_factor = lr_g_factor |
| 150 | self.trainable_ae_params = trainable_ae_params |
| 151 | if self.trainable_ae_params is not None: |
| 152 | self.ae_optimizer_args = default( |
| 153 | ae_optimizer_args, |
| 154 | [{} for _ in range(len(self.trainable_ae_params))], |
| 155 | ) |
| 156 | assert len(self.ae_optimizer_args) == len(self.trainable_ae_params) |
| 157 | else: |
| 158 | self.ae_optimizer_args = [{}] # makes type consitent |
| 159 | |
| 160 | self.trainable_disc_params = trainable_disc_params |
| 161 | if self.trainable_disc_params is not None: |
| 162 | self.disc_optimizer_args = default( |
| 163 | disc_optimizer_args, |
| 164 | [{} for _ in range(len(self.trainable_disc_params))], |
| 165 | ) |
| 166 | assert len(self.disc_optimizer_args) == len(self.trainable_disc_params) |
| 167 | else: |
| 168 | self.disc_optimizer_args = [{}] # makes type consitent |
| 169 | |
| 170 | if ckpt_path is not None: |
| 171 | assert ckpt_engine is None, "Can't set ckpt_engine and ckpt_path" |
| 172 | logpy.warn("Checkpoint path is deprecated, use `checkpoint_egnine` instead") |
| 173 | self.apply_ckpt(default(ckpt_path, ckpt_engine)) |
| 174 | self.additional_decode_keys = set(default(additional_decode_keys, [])) |
| 175 | |
| 176 | def get_input(self, batch: Dict) -> torch.Tensor: |