(self)
| 345 | #raise Exception |
| 346 | |
| 347 | def test(self): |
| 348 | if self.trainer is None: |
| 349 | self.lob_lightning_module = LOBLightningModule( |
| 350 | self.model, |
| 351 | experiment_id=self.experiment_id, |
| 352 | learning_rate=self.learning_rate, |
| 353 | general_hyperparameters=self.general_hyperparameters, |
| 354 | model_hyperparameters=self.model_hyperparameters, |
| 355 | ) |
| 356 | self.trainer = pl.Trainer() |
| 357 | try: |
| 358 | best_model = self.lob_lightning_module.load_from_checkpoint( |
| 359 | checkpoint_path=f"{logger.find_save_path(self.experiment_id)}/best_val_model.ckpt", |
| 360 | model=self.model, |
| 361 | experiment_id=self.experiment_id, |
| 362 | learning_rate=self.learning_rate, |
| 363 | general_hyperparameters=self.general_hyperparameters, |
| 364 | model_hyperparameters=self.model_hyperparameters, |
| 365 | ) |
| 366 | except: |
| 367 | best_model = self.lob_lightning_module.load_from_checkpoint( |
| 368 | checkpoint_path=f"{logger.find_save_path(self.experiment_id)}/best_val_model.ckpt", |
| 369 | map_location=torch.device('cpu'), |
| 370 | model=self.model, |
| 371 | experiment_id=self.experiment_id, |
| 372 | learning_rate=self.learning_rate, |
| 373 | general_hyperparameters=self.general_hyperparameters, |
| 374 | model_hyperparameters=self.model_hyperparameters, |
| 375 | ) |
| 376 | self.trainer.test(best_model, dataloaders=self.test_loader) |
| 377 | else: |
| 378 | best_model_path = ( |
| 379 | f"{logger.find_save_path(self.experiment_id)}/best_val_model.ckpt" |
| 380 | ) |
| 381 | self.trainer.test(ckpt_path=best_model_path, dataloaders=self.test_loader) |
no test coverage detected