(
self,
*args,
encoder_config: Dict,
decoder_config: Dict,
loss_config: Dict,
regularizer_config: Dict,
optimizer_config: Union[Dict, None] = None,
lr_g_factor: float = 1.0,
trainable_ae_params: Optional[List[List[str]]] = None,
ae_optimizer_args: Optional[List[dict]] = None,
trainable_disc_params: Optional[List[List[str]]] = None,
disc_optimizer_args: Optional[List[dict]] = None,
disc_start_iter: int = 0,
diff_boost_factor: float = 3.0,
ckpt_engine: Union[None, str, dict] = None,
ckpt_path: Optional[str] = None,
additional_decode_keys: Optional[List[str]] = None,
**kwargs,
)
| 133 | """ |
| 134 | |
| 135 | def __init__( |
| 136 | self, |
| 137 | *args, |
| 138 | encoder_config: Dict, |
| 139 | decoder_config: Dict, |
| 140 | loss_config: Dict, |
| 141 | regularizer_config: Dict, |
| 142 | optimizer_config: Union[Dict, None] = None, |
| 143 | lr_g_factor: float = 1.0, |
| 144 | trainable_ae_params: Optional[List[List[str]]] = None, |
| 145 | ae_optimizer_args: Optional[List[dict]] = None, |
| 146 | trainable_disc_params: Optional[List[List[str]]] = None, |
| 147 | disc_optimizer_args: Optional[List[dict]] = None, |
| 148 | disc_start_iter: int = 0, |
| 149 | diff_boost_factor: float = 3.0, |
| 150 | ckpt_engine: Union[None, str, dict] = None, |
| 151 | ckpt_path: Optional[str] = None, |
| 152 | additional_decode_keys: Optional[List[str]] = None, |
| 153 | **kwargs, |
| 154 | ): |
| 155 | super().__init__(*args, **kwargs) |
| 156 | self.automatic_optimization = False # pytorch lightning |
| 157 | |
| 158 | self.encoder = instantiate_from_config(encoder_config) |
| 159 | self.decoder = instantiate_from_config(decoder_config) |
| 160 | self.loss = instantiate_from_config(loss_config) |
| 161 | self.regularization = instantiate_from_config(regularizer_config) |
| 162 | self.optimizer_config = default(optimizer_config, {"target": "torch.optim.Adam"}) |
| 163 | self.diff_boost_factor = diff_boost_factor |
| 164 | self.disc_start_iter = disc_start_iter |
| 165 | self.lr_g_factor = lr_g_factor |
| 166 | self.trainable_ae_params = trainable_ae_params |
| 167 | if self.trainable_ae_params is not None: |
| 168 | self.ae_optimizer_args = default( |
| 169 | ae_optimizer_args, |
| 170 | [{} for _ in range(len(self.trainable_ae_params))], |
| 171 | ) |
| 172 | assert len(self.ae_optimizer_args) == len(self.trainable_ae_params) |
| 173 | else: |
| 174 | self.ae_optimizer_args = [{}] # makes type consitent |
| 175 | |
| 176 | self.trainable_disc_params = trainable_disc_params |
| 177 | if self.trainable_disc_params is not None: |
| 178 | self.disc_optimizer_args = default( |
| 179 | disc_optimizer_args, |
| 180 | [{} for _ in range(len(self.trainable_disc_params))], |
| 181 | ) |
| 182 | assert len(self.disc_optimizer_args) == len(self.trainable_disc_params) |
| 183 | else: |
| 184 | self.disc_optimizer_args = [{}] # makes type consitent |
| 185 | |
| 186 | if ckpt_path is not None: |
| 187 | assert ckpt_engine is None, "Can't set ckpt_engine and ckpt_path" |
| 188 | logpy.warn("Checkpoint path is deprecated, use `checkpoint_egnine` instead") |
| 189 | self.apply_ckpt(default(ckpt_path, ckpt_engine)) |
| 190 | self.additional_decode_keys = set(default(additional_decode_keys, [])) |
| 191 | |
| 192 | def get_input(self, batch: Dict) -> torch.Tensor: |
nothing calls this directly
no test coverage detected