mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
+1
-1
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user