| 614 | return out, None |
| 615 | |
| 616 | class LNNP(LightningModule): |
| 617 | def __init__(self, hparams, prior_model=None, mean=None, std=None): |
| 618 | super(LNNP, self).__init__() |
| 619 | |
| 620 | self.save_hyperparameters(hparams) |
| 621 | |
| 622 | if self.hparams.load_model: |
| 623 | self.model = load_model(self.hparams.load_model, args=self.hparams) |
| 624 | else: |
| 625 | self.model = create_model(self.hparams, prior_model, mean, std) |
| 626 | |
| 627 | self._reset_losses_dict() |
| 628 | self._reset_ema_dict() |
| 629 | self._reset_inference_results() |
| 630 | |
| 631 | def configure_optimizers(self): |
| 632 | optimizer = AdamW( |
| 633 | self.model.parameters(), |
| 634 | lr=self.hparams.lr, |
| 635 | weight_decay=self.hparams.weight_decay, |
| 636 | ) |
| 637 | scheduler = ReduceLROnPlateau( |
| 638 | optimizer, |
| 639 | "min", |
| 640 | factor=self.hparams.lr_factor, |
| 641 | patience=self.hparams.lr_patience, |
| 642 | min_lr=self.hparams.lr_min, |
| 643 | ) |
| 644 | lr_scheduler = { |
| 645 | "scheduler": scheduler, |
| 646 | "monitor": "val_loss", |
| 647 | "interval": "epoch", |
| 648 | "frequency": 1, |
| 649 | } |
| 650 | return [optimizer], [lr_scheduler] |
| 651 | |
| 652 | def forward(self, data): |
| 653 | return self.model(data) |
| 654 | |
| 655 | def training_step(self, batch, batch_idx): |
| 656 | loss_fn = mse_loss if self.hparams.loss_type == 'MSE' else l1_loss |
| 657 | |
| 658 | return self.step(batch, loss_fn, "train") |
| 659 | |
| 660 | def validation_step(self, batch, batch_idx, *args): |
| 661 | if len(args) == 0 or (len(args) > 0 and args[0] == 0): |
| 662 | # validation step |
| 663 | return self.step(batch, mse_loss, "val") |
| 664 | # test step |
| 665 | return self.step(batch, l1_loss, "test") |
| 666 | |
| 667 | def test_step(self, batch, batch_idx): |
| 668 | return self.step(batch, l1_loss, "test") |
| 669 | |
| 670 | def step(self, batch, loss_fn, stage): |
| 671 | with torch.set_grad_enabled(stage == "train" or self.hparams.derivative): |
| 672 | pred, deriv = self(batch) |
| 673 | if stage == "test": |