From c869dd8b8f6301f3726df84535a3da4e9acf04ec Mon Sep 17 00:00:00 2001 From: Jirka Borovec Date: Mon, 30 Mar 2020 18:14:27 +0200 Subject: [PATCH] make evaluate private (#1260) * make evaluate private * changelog --- CHANGELOG.md | 2 +- pytorch_lightning/trainer/evaluation_loop.py | 4 ++-- pytorch_lightning/trainer/trainer.py | 8 ++++---- tests/test_deprecated.py | 4 ++-- 4 files changed, 9 insertions(+), 9 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 90c9f495..4410702f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -20,7 +20,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). ### Changed -- +- Made `evalaute` method private >> `Trainer._evaluate(...)`. ([#1260](https://github.com/PyTorchLightning/pytorch-lightning/pull/1260)) ### Deprecated diff --git a/pytorch_lightning/trainer/evaluation_loop.py b/pytorch_lightning/trainer/evaluation_loop.py index bfbca4c2..4e9aed6d 100644 --- a/pytorch_lightning/trainer/evaluation_loop.py +++ b/pytorch_lightning/trainer/evaluation_loop.py @@ -217,7 +217,7 @@ class TrainerEvaluationLoopMixin(ABC): def reset_val_dataloader(self, *args): """Warning: this is just empty shell for code implemented in other class.""" - def evaluate(self, model: LightningModule, dataloaders, max_batches: int, test_mode: bool = False): + def _evaluate(self, model: LightningModule, dataloaders, max_batches: int, test_mode: bool = False): """Run evaluation code. Args: @@ -365,7 +365,7 @@ class TrainerEvaluationLoopMixin(ABC): setattr(self, f'{"test" if test_mode else "val"}_progress_bar', pbar) # run evaluation - eval_results = self.evaluate(self.model, dataloaders, max_batches, test_mode) + eval_results = self._evaluate(self.model, dataloaders, max_batches, test_mode) _, prog_bar_metrics, log_metrics, callback_metrics, _ = self.process_output( eval_results) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 9994af6c..2dc6c765 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -893,10 +893,10 @@ class Trainer( # dummy validation progress bar self.val_progress_bar = tqdm(disable=True) - eval_results = self.evaluate(model, - self.val_dataloaders, - self.num_sanity_val_steps, - False) + eval_results = self._evaluate(model, + self.val_dataloaders, + self.num_sanity_val_steps, + False) _, _, _, callback_metrics, _ = self.process_output(eval_results) # close progress bars diff --git a/tests/test_deprecated.py b/tests/test_deprecated.py index ddaae354..a3b087c7 100644 --- a/tests/test_deprecated.py +++ b/tests/test_deprecated.py @@ -95,7 +95,7 @@ def test_tbd_remove_in_v1_0_0_model_hooks(): trainer = Trainer(logger=False) # TODO: why `dataloder` is required if it is not used - result = trainer.evaluate(model, dataloaders=[[None]], max_batches=1) + result = trainer._evaluate(model, dataloaders=[[None]], max_batches=1) assert result == {'val_loss': 0.6} model = ModelVer0_7(hparams) @@ -106,5 +106,5 @@ def test_tbd_remove_in_v1_0_0_model_hooks(): trainer = Trainer(logger=False) # TODO: why `dataloder` is required if it is not used - result = trainer.evaluate(model, dataloaders=[[None]], max_batches=1) + result = trainer._evaluate(model, dataloaders=[[None]], max_batches=1) assert result == {'val_loss': 0.7}