make evaluate private (#1260)

* make evaluate private

* changelog
This commit is contained in:
Jirka Borovec
2020-03-30 12:14:27 -04:00
committed by GitHub
parent 6dfe9951e1
commit c869dd8b8f
4 changed files with 9 additions and 9 deletions
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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)
+4 -4
View File
@@ -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
+2 -2
View File
@@ -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}