(self, batch: dict, batch_idx: int, optimizer_idx: int = 0)
| 224 | return z, dec, reg_log |
| 225 | |
| 226 | def inner_training_step(self, batch: dict, batch_idx: int, optimizer_idx: int = 0) -> torch.Tensor: |
| 227 | x = self.get_input(batch) |
| 228 | additional_decode_kwargs = {key: batch[key] for key in self.additional_decode_keys.intersection(batch)} |
| 229 | z, xrec, regularization_log = self(x, **additional_decode_kwargs) |
| 230 | if hasattr(self.loss, "forward_keys"): |
| 231 | extra_info = { |
| 232 | "z": z, |
| 233 | "optimizer_idx": optimizer_idx, |
| 234 | "global_step": self.global_step, |
| 235 | "last_layer": self.get_last_layer(), |
| 236 | "split": "train", |
| 237 | "regularization_log": regularization_log, |
| 238 | "autoencoder": self, |
| 239 | } |
| 240 | extra_info = {k: extra_info[k] for k in self.loss.forward_keys} |
| 241 | else: |
| 242 | extra_info = dict() |
| 243 | |
| 244 | if optimizer_idx == 0: |
| 245 | # autoencode |
| 246 | out_loss = self.loss(x, xrec, **extra_info) |
| 247 | if isinstance(out_loss, tuple): |
| 248 | aeloss, log_dict_ae = out_loss |
| 249 | else: |
| 250 | # simple loss function |
| 251 | aeloss = out_loss |
| 252 | log_dict_ae = {"train/loss/rec": aeloss.detach()} |
| 253 | |
| 254 | self.log_dict( |
| 255 | log_dict_ae, |
| 256 | prog_bar=False, |
| 257 | logger=True, |
| 258 | on_step=True, |
| 259 | on_epoch=True, |
| 260 | sync_dist=False, |
| 261 | ) |
| 262 | self.log( |
| 263 | "loss", |
| 264 | aeloss.mean().detach(), |
| 265 | prog_bar=True, |
| 266 | logger=False, |
| 267 | on_epoch=False, |
| 268 | on_step=True, |
| 269 | ) |
| 270 | return aeloss |
| 271 | elif optimizer_idx == 1: |
| 272 | # discriminator |
| 273 | discloss, log_dict_disc = self.loss(x, xrec, **extra_info) |
| 274 | # -> discriminator always needs to return a tuple |
| 275 | self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=True) |
| 276 | return discloss |
| 277 | else: |
| 278 | raise NotImplementedError(f"Unknown optimizer {optimizer_idx}") |
| 279 | |
| 280 | def training_step(self, batch: dict, batch_idx: int): |
| 281 | opts = self.optimizers() |
no test coverage detected