Tests: refactor cleanup (#1744)

* wip

* cleaning

* optim imports

* -

* default hparams

* fix restore

* fix imports
This commit is contained in:
Jirka Borovec
2020-05-10 13:15:28 -04:00
committed by GitHub
parent 4970927ec8
commit 134eb61e1a
41 changed files with 187 additions and 1150 deletions
+2 -59
View File
@@ -1,61 +1,4 @@
"""Models for testing."""
import torch
from tests.base.eval_model_template import EvalModelTemplate
from tests.base.mixins import (
LightEmptyTestStep,
LightValidationStepMixin,
LightValidationMixin,
LightValidationStepMultipleDataloadersMixin,
LightValidationMultipleDataloadersMixin,
LightTestStepMixin,
LightTestMixin,
LightTestStepMultipleDataloadersMixin,
LightTestMultipleDataloadersMixin,
LightTestFitSingleTestDataloadersMixin,
LightTestFitMultipleTestDataloadersMixin,
LightValStepFitSingleDataloaderMixin,
LightValStepFitMultipleDataloadersMixin,
LightTrainDataloader,
LightValidationDataloader,
LightTestDataloader,
LightInfTrainDataloader,
LightInfValDataloader,
LightInfTestDataloader,
LightTestOptimizerWithSchedulingMixin,
LightTestMultipleOptimizersWithSchedulingMixin,
LightTestOptimizersWithMixedSchedulingMixin,
LightTestReduceLROnPlateauMixin,
LightTestNoneOptimizerMixin,
LightZeroLenDataloader
)
from tests.base.models import TestModelBase, DictHparamsModel
class LightningTestModel(LightTrainDataloader,
LightValidationMixin,
LightTestMixin,
TestModelBase):
"""Most common test case. Validation and test dataloaders."""
def on_training_metrics(self, logs):
logs['some_tensor_to_test'] = torch.rand(1)
class LightningTestModelWithoutHyperparametersArg(LightningTestModel):
"""Without hparams argument in constructor """
def __init__(self):
import tests.base.utils as tutils
# the user loads the hparams in some other way
hparams = tutils.get_default_hparams()
super().__init__(hparams)
class LightningTestModelWithUnusedHyperparametersArg(LightningTestModelWithoutHyperparametersArg):
"""It has hparams argument in constructor but is not used."""
def __init__(self, hparams):
super().__init__()
from tests.base.datasets import TrialMNIST
from tests.base.model_template import EvalModelTemplate
+23
View File
@@ -0,0 +1,23 @@
"""Custom dataloaders for testing"""
class CustomInfDataloader:
def __init__(self, dataloader):
self.dataloader = dataloader
self.iter = iter(dataloader)
self.count = 0
def __iter__(self):
self.count = 0
return self
def __next__(self):
if self.count >= 50:
raise StopIteration
self.count = self.count + 1
try:
return next(self.iter)
except StopIteration:
self.iter = iter(self.dataloader)
return next(self.iter)
-51
View File
@@ -1,51 +0,0 @@
import torch
from torch.nn import functional as F
from torch.utils.data import DataLoader
import pytorch_lightning as pl
from tests.base.datasets import TrialMNIST
# from test_models import assert_ok_test_acc, load_model, \
# clear_save_dir, get_default_logger, get_default_hparams, init_save_dir, \
# init_checkpoint_callback, reset_seed, set_random_master_port
class CoolModel(pl.LightningModule):
def __init(self):
super().__init__()
# not the best model...
self.l1 = torch.nn.Linear(28 * 28, 10)
def forward(self, x):
return torch.relu(self.l1(x))
def my_loss(self, y_hat, y):
return F.cross_entropy(y_hat, y)
def training_step(self, batch, batch_idx):
x, y = batch
y_hat = self(x)
return {'training_loss': self.my_loss(y_hat, y)}
def validation_step(self, batch, batch_idx):
x, y = batch
y_hat = self(x)
return {'val_loss': self.my_loss(y_hat, y)}
def validation_epoch_end(self, outputs):
avg_loss = torch.stack([x for x in outputs['val_loss']]).mean()
return avg_loss
def configure_optimizers(self):
return [torch.optim.Adam(self.parameters(), lr=0.02)]
def train_dataloader(self):
return DataLoader(TrialMNIST(train=True, num_samples=100), batch_size=16)
def val_dataloader(self):
return DataLoader(TrialMNIST(train=False, num_samples=50), batch_size=16)
def test_dataloader(self):
return DataLoader(TrialMNIST(train=False, num_samples=50), batch_size=16)
-718
View File
@@ -1,718 +0,0 @@
from collections import OrderedDict
import torch
from torch import optim
class LightValidationStepMixin:
"""
Add val_dataloader and validation_step methods for the case
when val_dataloader returns a single dataloader
"""
def val_dataloader(self):
return self._dataloader(train=False)
def validation_step(self, batch, batch_idx, *args, **kwargs):
"""Lightning calls this inside the validation loop."""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_val = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
val_acc = torch.tensor(val_acc)
if self.on_gpu:
val_acc = val_acc.cuda(loss_val.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_val = loss_val.unsqueeze(0)
val_acc = val_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
})
return output
if batch_idx % 2 == 0:
return val_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
'test_dic': {'val_loss_a': loss_val}
})
return output
class LightValidationMixin(LightValidationStepMixin):
"""
Add val_dataloader, validation_step, and validation_end methods for the case
when val_dataloader returns a single dataloader
"""
def validation_epoch_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
Args:
outputs: list of individual outputs of each validation step
"""
# if returned a scalar from validation_step, outputs is a list of tensor scalars
# we return just the average in this case (if we want)
# return torch.stack(outputs).mean()
val_loss_mean = 0
val_acc_mean = 0
for output in outputs:
val_loss = _get_output_metric(output, 'val_loss')
# reduce manually when using dp
if self.trainer.use_dp or self.trainer.use_ddp2:
val_loss = torch.mean(val_loss)
val_loss_mean += val_loss
# reduce manually when using dp
val_acc = _get_output_metric(output, 'val_acc')
if self.trainer.use_dp or self.trainer.use_ddp2:
val_acc = torch.mean(val_acc)
val_acc_mean += val_acc
val_loss_mean /= len(outputs)
val_acc_mean /= len(outputs)
metrics_dict = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()}
results = {'progress_bar': metrics_dict, 'log': metrics_dict}
return results
class LightValidationStepMultipleDataloadersMixin:
"""
Add val_dataloader and validation_step methods for the case
when val_dataloader returns multiple dataloaders
"""
def val_dataloader(self):
return [self._dataloader(train=False), self._dataloader(train=False)]
def validation_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
:return:
"""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_val = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
val_acc = torch.tensor(val_acc)
if self.on_gpu:
val_acc = val_acc.cuda(loss_val.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_val = loss_val.unsqueeze(0)
val_acc = val_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
})
return output
if batch_idx % 2 == 0:
return val_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
'test_dic': {'val_loss_a': loss_val}
})
return output
if batch_idx % 5 == 0:
output = OrderedDict({
f'val_loss_{dataloader_idx}': loss_val,
f'val_acc_{dataloader_idx}': val_acc,
})
return output
class LightValidationMultipleDataloadersMixin(LightValidationStepMultipleDataloadersMixin):
"""
Add val_dataloader, validation_step, and validation_end methods for the case
when val_dataloader returns multiple dataloaders
"""
def validation_epoch_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
:param outputs: list of individual outputs of each validation step
:return:
"""
# if returned a scalar from validation_step, outputs is a list of tensor scalars
# we return just the average in this case (if we want)
# return torch.stack(outputs).mean()
val_loss_mean = 0
val_acc_mean = 0
i = 0
for dl_output in outputs:
for output in dl_output:
val_loss = output['val_loss']
# reduce manually when using dp
if self.trainer.use_dp:
val_loss = torch.mean(val_loss)
val_loss_mean += val_loss
# reduce manually when using dp
val_acc = output['val_acc']
if self.trainer.use_dp:
val_acc = torch.mean(val_acc)
val_acc_mean += val_acc
i += 1
val_loss_mean /= i
val_acc_mean /= i
tqdm_dict = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()}
result = {'progress_bar': tqdm_dict}
return result
class LightTrainDataloader:
"""Simple train dataloader."""
def train_dataloader(self):
return self._dataloader(train=True)
class LightValidationDataloader:
"""Simple validation dataloader."""
def val_dataloader(self):
return self._dataloader(train=False)
class LightTestDataloader:
"""Simple test dataloader."""
def test_dataloader(self):
return self._dataloader(train=False)
class CustomInfDataloader:
def __init__(self, dataloader):
self.dataloader = dataloader
self.iter = iter(dataloader)
self.count = 0
def __iter__(self):
self.count = 0
return self
def __next__(self):
if self.count >= 50:
raise StopIteration
self.count = self.count + 1
try:
return next(self.iter)
except StopIteration:
self.iter = iter(self.dataloader)
return next(self.iter)
class LightInfTrainDataloader:
"""Simple test dataloader."""
def train_dataloader(self):
return CustomInfDataloader(self._dataloader(train=True))
class LightInfValDataloader:
"""Simple test dataloader."""
def val_dataloader(self):
return CustomInfDataloader(self._dataloader(train=False))
class LightInfTestDataloader:
"""Simple test dataloader."""
def test_dataloader(self):
return CustomInfDataloader(self._dataloader(train=False))
class LightZeroLenDataloader:
""" Simple dataloader that has zero length. """
def train_dataloader(self):
dataloader = self._dataloader(train=True)
dataloader.dataset.data = dataloader.dataset.data[:0]
dataloader.dataset.targets = dataloader.dataset.targets[:0]
return dataloader
class LightEmptyTestStep:
"""Empty test step."""
def test_step(self, *args, **kwargs):
return dict()
class LightTestStepMixin(LightTestDataloader):
"""Test step mixin."""
def test_step(self, batch, batch_idx, *args, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
:return:
"""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_test = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
test_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
test_acc = torch.tensor(test_acc)
if self.on_gpu:
test_acc = test_acc.cuda(loss_test.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_test = loss_test.unsqueeze(0)
test_acc = test_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
})
return output
if batch_idx % 2 == 0:
return test_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
'test_dic': {'test_loss_a': loss_test}
})
return output
class LightTestMixin(LightTestStepMixin):
"""Ritch test mixin."""
def test_epoch_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
:param outputs: list of individual outputs of each validation step
:return:
"""
# if returned a scalar from test_step, outputs is a list of tensor scalars
# we return just the average in this case (if we want)
# return torch.stack(outputs).mean()
test_loss_mean = 0
test_acc_mean = 0
for output in outputs:
test_loss = _get_output_metric(output, 'test_loss')
# reduce manually when using dp
if self.trainer.use_dp:
test_loss = torch.mean(test_loss)
test_loss_mean += test_loss
# reduce manually when using dp
test_acc = _get_output_metric(output, 'test_acc')
if self.trainer.use_dp:
test_acc = torch.mean(test_acc)
test_acc_mean += test_acc
test_loss_mean /= len(outputs)
test_acc_mean /= len(outputs)
metrics_dict = {'test_loss': test_loss_mean.item(), 'test_acc': test_acc_mean.item()}
result = {'progress_bar': metrics_dict, 'log': metrics_dict}
return result
class LightTestStepMultipleDataloadersMixin:
"""Test step multiple dataloaders mixin."""
def test_dataloader(self):
return [self._dataloader(train=False), self._dataloader(train=False)]
def test_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
:return:
"""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_test = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
test_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
test_acc = torch.tensor(test_acc)
if self.on_gpu:
test_acc = test_acc.cuda(loss_test.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_test = loss_test.unsqueeze(0)
test_acc = test_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
})
return output
if batch_idx % 2 == 0:
return test_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
'test_dic': {'test_loss_a': loss_test}
})
return output
if batch_idx % 5 == 0:
output = OrderedDict({
f'test_loss_{dataloader_idx}': loss_test,
f'test_acc_{dataloader_idx}': test_acc,
})
return output
class LightTestFitSingleTestDataloadersMixin:
"""Test fit single test dataloaders mixin."""
def test_dataloader(self):
return self._dataloader(train=False)
def test_step(self, batch, batch_idx, *args, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
:return:
"""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_test = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
test_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
test_acc = torch.tensor(test_acc)
if self.on_gpu:
test_acc = test_acc.cuda(loss_test.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_test = loss_test.unsqueeze(0)
test_acc = test_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
})
return output
if batch_idx % 2 == 0:
return test_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
'test_dic': {'test_loss_a': loss_test}
})
return output
class LightTestFitMultipleTestDataloadersMixin:
"""Test fit multiple test dataloaders mixin."""
def test_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
:return:
"""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_test = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
test_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
test_acc = torch.tensor(test_acc)
if self.on_gpu:
test_acc = test_acc.cuda(loss_test.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_test = loss_test.unsqueeze(0)
test_acc = test_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
})
return output
if batch_idx % 2 == 0:
return test_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'test_loss': loss_test,
'test_acc': test_acc,
'test_dic': {'test_loss_a': loss_test}
})
return output
if batch_idx % 5 == 0:
output = OrderedDict({
f'test_loss_{dataloader_idx}': loss_test,
f'test_acc_{dataloader_idx}': test_acc,
})
return output
class LightValStepFitSingleDataloaderMixin:
def validation_step(self, batch, batch_idx, *args, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
:return:
"""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_val = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
val_acc = torch.tensor(val_acc)
if self.on_gpu:
val_acc = val_acc.cuda(loss_val.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_val = loss_val.unsqueeze(0)
val_acc = val_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
})
return output
if batch_idx % 2 == 0:
return val_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
'test_dic': {'val_loss_a': loss_val}
})
return output
class LightValStepFitMultipleDataloadersMixin:
def validation_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
:return:
"""
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
loss_val = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
val_acc = torch.tensor(val_acc)
if self.on_gpu:
val_acc = val_acc.cuda(loss_val.device.index)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_val = loss_val.unsqueeze(0)
val_acc = val_acc.unsqueeze(0)
# alternate possible outputs to test
if batch_idx % 1 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
})
return output
if batch_idx % 2 == 0:
return val_acc
if batch_idx % 3 == 0:
output = OrderedDict({
'val_loss': loss_val,
'val_acc': val_acc,
'test_dic': {'val_loss_a': loss_val}
})
return output
if batch_idx % 5 == 0:
output = OrderedDict({
f'val_loss_{dataloader_idx}': loss_val,
f'val_acc_{dataloader_idx}': val_acc,
})
return output
class LightTestMultipleDataloadersMixin(LightTestStepMultipleDataloadersMixin):
def test_epoch_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
:param outputs: list of individual outputs of each validation step
:return:
"""
# if returned a scalar from test_step, outputs is a list of tensor scalars
# we return just the average in this case (if we want)
# return torch.stack(outputs).mean()
test_loss_mean = 0
test_acc_mean = 0
i = 0
for dl_output in outputs:
for output in dl_output:
test_loss = output['test_loss']
# reduce manually when using dp
if self.trainer.use_dp:
test_loss = torch.mean(test_loss)
test_loss_mean += test_loss
# reduce manually when using dp
test_acc = output['test_acc']
if self.trainer.use_dp:
test_acc = torch.mean(test_acc)
test_acc_mean += test_acc
i += 1
test_loss_mean /= i
test_acc_mean /= i
tqdm_dict = {'test_loss': test_loss_mean.item(), 'test_acc': test_acc_mean.item()}
result = {'progress_bar': tqdm_dict}
return result
class LightTestOptimizerWithSchedulingMixin:
def configure_optimizers(self):
if self.hparams.optimizer_name == 'lbfgs':
optimizer = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
else:
optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
lr_scheduler = optim.lr_scheduler.StepLR(optimizer, 1, gamma=0.1)
return [optimizer], [lr_scheduler]
class LightTestMultipleOptimizersWithSchedulingMixin:
def configure_optimizers(self):
if self.hparams.optimizer_name == 'lbfgs':
optimizer1 = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
optimizer2 = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
else:
optimizer1 = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
optimizer2 = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
lr_scheduler1 = optim.lr_scheduler.StepLR(optimizer1, 1, gamma=0.1)
lr_scheduler2 = optim.lr_scheduler.StepLR(optimizer2, 1, gamma=0.1)
return [optimizer1, optimizer2], [lr_scheduler1, lr_scheduler2]
class LightTestOptimizersWithMixedSchedulingMixin:
def configure_optimizers(self):
if self.hparams.optimizer_name == 'lbfgs':
optimizer1 = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
optimizer2 = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
else:
optimizer1 = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
optimizer2 = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
lr_scheduler1 = optim.lr_scheduler.StepLR(optimizer1, 4, gamma=0.1)
lr_scheduler2 = optim.lr_scheduler.StepLR(optimizer2, 1, gamma=0.1)
return [optimizer1, optimizer2], \
[{'scheduler': lr_scheduler1, 'interval': 'step'}, lr_scheduler2]
class LightTestReduceLROnPlateauMixin:
def configure_optimizers(self):
if self.hparams.optimizer_name == 'lbfgs':
optimizer = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
else:
optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
lr_scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer)
return [optimizer], [lr_scheduler]
class LightTestNoneOptimizerMixin:
def configure_optimizers(self):
return None
def _get_output_metric(output, name):
if isinstance(output, dict):
val = output[name]
else: # if it is 2level deep -> per dataloader and per batch
val = sum(out[name] for out in output) / len(output)
return val
@@ -5,17 +5,17 @@ import torch.nn as nn
import torch.nn.functional as F
from pytorch_lightning.core.lightning import LightningModule
from tests.base.datasets import TrialMNIST
from tests.base.eval_model_optimizers import ConfigureOptimizersPool
from tests.base.eval_model_test_dataloaders import TestDataloaderVariations
from tests.base.eval_model_test_epoch_ends import TestEpochEndVariations
from tests.base.eval_model_test_steps import TestStepVariations
from tests.base.eval_model_train_dataloaders import TrainDataloaderVariations
from tests.base.eval_model_train_steps import TrainingStepVariations
from tests.base.eval_model_utils import ModelTemplateUtils, ModelTemplateData
from tests.base.eval_model_valid_dataloaders import ValDataloaderVariations
from tests.base.eval_model_valid_epoch_ends import ValidationEpochEndVariations
from tests.base.eval_model_valid_steps import ValidationStepVariations
from tests.base.datasets import TrialMNIST, PATH_DATASETS
from tests.base.model_optimizers import ConfigureOptimizersPool
from tests.base.model_test_dataloaders import TestDataloaderVariations
from tests.base.model_test_epoch_ends import TestEpochEndVariations
from tests.base.model_test_steps import TestStepVariations
from tests.base.model_train_dataloaders import TrainDataloaderVariations
from tests.base.model_train_steps import TrainingStepVariations
from tests.base.model_utilities import ModelTemplateUtils, ModelTemplateData
from tests.base.model_valid_dataloaders import ValDataloaderVariations
from tests.base.model_valid_epoch_ends import ValidationEpochEndVariations
from tests.base.model_valid_steps import ValidationStepVariations
class EvalModelTemplate(
@@ -35,8 +35,10 @@ class EvalModelTemplate(
"""
This template houses all combinations of model configurations we want to test
"""
def __init__(self, hparams: object) -> object:
def __init__(self, hparams: object = None) -> object:
"""Pass in parsed HyperOptArgumentParser to the model."""
if hparams is None:
hparams = EvalModelTemplate.get_default_hparams()
# init superclass
super().__init__()
self.hparams = Namespace(**hparams) if isinstance(hparams, dict) else hparams
@@ -81,3 +83,27 @@ class EvalModelTemplate(
def prepare_data(self):
_ = TrialMNIST(root=self.hparams.data_root, train=True, download=True)
@staticmethod
def get_default_hparams(continue_training: bool = False, hpc_exp_number: int = 0) -> Namespace:
args = dict(
drop_prob=0.2,
batch_size=32,
in_features=28 * 28,
learning_rate=0.001 * 8,
optimizer_name='adam',
data_root=PATH_DATASETS,
out_features=10,
hidden_dim=1000,
b1=0.5,
b2=0.999,
)
if continue_training:
args.update(
test_tube_do_checkpoint_load=True,
hpc_exp_number=hpc_exp_number,
)
hparams = Namespace(**args)
return hparams
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod
from tests.base.eval_model_utils import CustomInfDataloader
from tests.base.dataloaders import CustomInfDataloader
class TestDataloaderVariations(ABC):
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod
from tests.base.eval_model_utils import CustomInfDataloader
from tests.base.dataloaders import CustomInfDataloader
class TrainDataloaderVariations(ABC):
@@ -26,25 +26,3 @@ class ModelTemplateUtils:
else: # if it is 2level deep -> per dataloader and per batch
val = sum(out[name] for out in output) / len(output)
return val
class CustomInfDataloader:
def __init__(self, dataloader):
self.dataloader = dataloader
self.iter = iter(dataloader)
self.count = 0
def __iter__(self):
self.count = 0
return self
def __next__(self):
if self.count >= 50:
raise StopIteration
self.count = self.count + 1
try:
return next(self.iter)
except StopIteration:
self.iter = iter(self.dataloader)
return next(self.iter)
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod
from tests.base.eval_model_utils import CustomInfDataloader
from tests.base.dataloaders import CustomInfDataloader
class ValDataloaderVariations(ABC):
-139
View File
@@ -1,14 +1,11 @@
from collections import OrderedDict
from typing import Dict
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import optim
from torch.utils.data import DataLoader
from tests.base import EvalModelTemplate
from tests.base.datasets import TrialMNIST
try:
@@ -20,142 +17,6 @@ except ImportError:
from pytorch_lightning.core.lightning import LightningModule
class DictHparamsModel(LightningModule):
def __init__(self, hparams: Dict):
super().__init__()
self.hparams = hparams
self.l1 = torch.nn.Linear(hparams.get('in_features'), hparams['out_features'])
def forward(self, x):
return torch.relu(self.l1(x.view(x.size(0), -1)))
def training_step(self, batch, batch_idx):
x, y = batch
y_hat = self(x)
return {'loss': F.cross_entropy(y_hat, y)}
def configure_optimizers(self):
return torch.optim.Adam(self.parameters(), lr=0.02)
def train_dataloader(self):
return DataLoader(TrialMNIST(train=True, download=True), batch_size=16)
class TestModelBase(LightningModule):
"""Base LightningModule for testing. Implements only the required interface."""
def __init__(self, hparams, force_remove_distributed_sampler: bool = False):
"""Pass in parsed HyperOptArgumentParser to the model."""
# init superclass
super().__init__()
self.hparams = hparams
self.batch_size = hparams.batch_size
# if you specify an example input, the summary will show input/output for each layer
self.example_input_array = torch.rand(5, 28 * 28)
# remove to test warning for dist sampler
self.force_remove_distributed_sampler = force_remove_distributed_sampler
# build model
self.__build_model()
# ---------------------
# MODEL SETUP
# ---------------------
def __build_model(self):
"""Layout model."""
self.c_d1 = nn.Linear(in_features=self.hparams.in_features,
out_features=self.hparams.hidden_dim)
self.c_d1_bn = nn.BatchNorm1d(self.hparams.hidden_dim)
self.c_d1_drop = nn.Dropout(self.hparams.drop_prob)
self.c_d2 = nn.Linear(in_features=self.hparams.hidden_dim,
out_features=self.hparams.out_features)
# ---------------------
# TRAINING
# ---------------------
def forward(self, x):
"""No special modification required for lightning, define as you normally would."""
x = self.c_d1(x)
x = torch.tanh(x)
x = self.c_d1_bn(x)
x = self.c_d1_drop(x)
x = self.c_d2(x)
logits = F.log_softmax(x, dim=1)
return logits
def loss(self, labels, logits):
nll = F.nll_loss(logits, labels)
return nll
def training_step(self, batch, batch_idx, optimizer_idx=None):
"""Lightning calls this inside the training loop"""
# forward pass
x, y = batch
x = x.view(x.size(0), -1)
y_hat = self(x)
# calculate loss
loss_val = self.loss(y, y_hat)
# in DP mode (default) make sure if result is scalar, there's another dim in the beginning
if self.trainer.use_dp:
loss_val = loss_val.unsqueeze(0)
# alternate possible outputs to test
if self.trainer.batch_idx % 1 == 0:
output = OrderedDict({
'loss': loss_val,
'progress_bar': {'some_val': loss_val * loss_val},
'log': {'train_some_val': loss_val * loss_val},
})
return output
if self.trainer.batch_idx % 2 == 0:
return loss_val
# ---------------------
# TRAINING SETUP
# ---------------------
def configure_optimizers(self):
"""
return whatever optimizers we want here.
:return: list of optimizers
"""
# try no scheduler for this model (testing purposes)
if self.hparams.optimizer_name == 'lbfgs':
optimizer = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
else:
optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)
return [optimizer], [scheduler]
def prepare_data(self):
_ = TrialMNIST(root=self.hparams.data_root, train=True, download=True)
def _dataloader(self, train):
# init data generators
dataset = TrialMNIST(root=self.hparams.data_root, train=train, download=True)
# when using multi-node we need to add the datasampler
batch_size = self.hparams.batch_size
loader = DataLoader(
dataset=dataset,
batch_size=batch_size,
shuffle=train
)
return loader
class Generator(nn.Module):
def __init__(self, latent_dim, img_shape):
super().__init__()
+3 -28
View File
@@ -9,8 +9,7 @@ from pytorch_lightning import Trainer
from pytorch_lightning.callbacks import ModelCheckpoint
from pytorch_lightning.loggers import TensorBoardLogger
from tests import TEMP_PATH, RANDOM_PORTS, RANDOM_SEEDS
from tests.base import LightningTestModel, EvalModelTemplate
from tests.base.datasets import PATH_DATASETS
from tests.base.model_template import EvalModelTemplate
def assert_speed_parity(pl_times, pt_times, num_epochs):
@@ -97,30 +96,6 @@ def run_model_test(trainer_options, model, on_gpu=True, version=None, with_hpc=T
trainer.hpc_load(save_dir, on_gpu=on_gpu)
def get_default_hparams(continue_training=False, hpc_exp_number=0):
args = {
'drop_prob': 0.2,
'batch_size': 32,
'in_features': 28 * 28,
'learning_rate': 0.001 * 8,
'optimizer_name': 'adam',
'data_root': PATH_DATASETS,
'out_features': 10,
'hidden_dim': 1000,
'b1': 0.5,
'b2': 0.999,
}
if continue_training:
args.update(
test_tube_do_checkpoint_load=True,
hpc_exp_number=hpc_exp_number,
)
hparams = Namespace(**args)
return hparams
def get_default_logger(save_dir, version=None):
# set up logger object without actually saving logs
logger = TensorBoardLogger(save_dir, name='lightning_logs', version=version)
@@ -148,7 +123,7 @@ def get_data_path(expt_logger, path_dir=None):
return path_expt
def load_model(logger, root_weights_dir, module_class=LightningTestModel, path_expt=None):
def load_model(logger, root_weights_dir, module_class=EvalModelTemplate, path_expt=None):
# load trained model
path_expt_dir = get_data_path(logger, path_dir=path_expt)
tags_path = os.path.join(path_expt_dir, TensorBoardLogger.NAME_CSV_TAGS)
@@ -166,7 +141,7 @@ def load_model(logger, root_weights_dir, module_class=LightningTestModel, path_e
return trained_model
def load_model_from_checkpoint(root_weights_dir, module_class=LightningTestModel):
def load_model_from_checkpoint(root_weights_dir, module_class=EvalModelTemplate):
# load trained model
checkpoints = [x for x in os.listdir(root_weights_dir) if '.ckpt' in x]
weights_dir = os.path.join(root_weights_dir, checkpoints[0])
+7 -6
View File
@@ -1,4 +1,5 @@
import pytest
import tests.base.utils as tutils
from pytorch_lightning import Callback
from pytorch_lightning import Trainer, LightningModule
@@ -11,7 +12,7 @@ from pathlib import Path
def test_trainer_callback_system(tmpdir):
"""Test the callback system."""
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
def _check_args(trainer, pl_module):
@@ -209,7 +210,7 @@ def test_early_stopping_no_val_step(tmpdir):
output.update({'my_train_metric': output['loss']}) # could be anything else
return output
model = CurrentModel(tutils.get_default_hparams())
model = CurrentModel()
model.validation_step = None
model.val_dataloader = None
@@ -245,7 +246,7 @@ def test_pickling(tmpdir):
def test_model_checkpoint_with_non_string_input(tmpdir, save_top_k):
""" Test that None in checkpoint callback is valid and that chkp_path is set correctly """
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
checkpoint = ModelCheckpoint(filepath=None, save_top_k=save_top_k)
@@ -267,7 +268,7 @@ def test_model_checkpoint_with_non_string_input(tmpdir, save_top_k):
def test_model_checkpoint_path(tmpdir, logger_version, expected):
"""Test that "version_" prefix is only added when logger's version is an integer"""
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
logger = TensorBoardLogger(str(tmpdir), version=logger_version)
trainer = Trainer(
@@ -286,7 +287,7 @@ def test_lr_logger_single_lr(tmpdir):
""" Test that learning rates are extracted and logged for single lr scheduler"""
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.configure_optimizers = model.configure_optimizers__single_scheduler
lr_logger = LearningRateLogger()
@@ -311,7 +312,7 @@ def test_lr_logger_multi_lrs(tmpdir):
""" Test that learning rates are extracted and logged for multi lr schedulers """
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.configure_optimizers = model.configure_optimizers__multiple_schedulers
lr_logger = LearningRateLogger()
+3 -3
View File
@@ -58,7 +58,7 @@ def test_progress_bar_misconfiguration():
def test_progress_bar_totals():
"""Test that the progress finishes with the correct total steps processed."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer = Trainer(
progress_bar_refresh_rate=1,
@@ -107,7 +107,7 @@ def test_progress_bar_totals():
def test_progress_bar_fast_dev_run():
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer = Trainer(
fast_dev_run=True,
@@ -140,7 +140,7 @@ def test_progress_bar_fast_dev_run():
def test_progress_bar_progress_refresh(refresh_rate):
"""Test that the three progress bars get correctly updated when using different refresh rates."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
class CurrentProgressBar(ProgressBar):
+1 -1
View File
@@ -35,7 +35,7 @@ def test_loggers_fit_test(tmpdir, monkeypatch, logger_class):
import atexit
monkeypatch.setattr(atexit, 'register', lambda _: None)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
class StoreHistoryLogger(logger_class):
def __init__(self, *args, **kwargs):
+4 -4
View File
@@ -60,8 +60,8 @@ class CustomLogger(LightningLoggerBase):
def test_custom_logger(tmpdir):
hparams = tutils.get_default_hparams()
model = EvalModelTemplate(tutils.get_default_hparams())
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
logger = CustomLogger()
@@ -79,7 +79,7 @@ def test_custom_logger(tmpdir):
def test_multiple_loggers(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
logger1 = CustomLogger()
@@ -139,7 +139,7 @@ def test_adding_step_key(tmpdir):
return decorated
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.validation_epoch_end = _validation_epoch_end
model.training_epoch_end = _training_epoch_end
trainer = Trainer(
+1 -1
View File
@@ -61,7 +61,7 @@ def test_neptune_additional_methods(neptune):
def test_neptune_leave_open_experiment_after_fit(tmpdir):
"""Verify that neptune experiment was closed after training"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
def _run_training(logger):
logger._experiment = MagicMock()
+1 -1
View File
@@ -8,7 +8,7 @@ from tests.base import EvalModelTemplate
def test_trains_logger(tmpdir):
"""Verify that basic functionality of TRAINS logger works."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
TrainsLogger.set_bypass_mode(True)
TrainsLogger.set_credentials(api_host='http://integration.trains.allegro.ai:8008',
files_host='http://integration.trains.allegro.ai:8081',
@@ -30,7 +30,7 @@ sys.path.insert(0, os.path.abspath(PATH_ROOT))
from pytorch_lightning import Trainer # noqa: E402
from pytorch_lightning.callbacks import ModelCheckpoint # noqa: E402
from tests.base import EvalModelTemplate # noqa: E402
from tests.base.utils import set_random_master_port, get_default_hparams, run_model_test # noqa: E402
from tests.base.utils import set_random_master_port, run_model_test # noqa: E402
parser = argparse.ArgumentParser()
@@ -45,7 +45,7 @@ def run_test_from_config(trainer_options):
ckpt_path = trainer_options['default_root_dir']
trainer_options.update(checkpoint_callback=ModelCheckpoint(ckpt_path))
model = EvalModelTemplate(get_default_hparams())
model = EvalModelTemplate(EvalModelTemplate.get_default_hparams())
run_model_test(trainer_options, model, on_gpu=args.on_gpu, version=0, with_hpc=False)
# Horovod should be initialized following training. If not, this will raise an exception.
+4 -4
View File
@@ -23,7 +23,7 @@ def test_amp_single_gpu(tmpdir, backend):
precision=16
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# tutils.run_model_test(trainer_options, model)
result = trainer.fit(model)
@@ -37,7 +37,7 @@ def test_amp_multi_gpu(tmpdir, backend):
"""Make sure DP/DDP + AMP work."""
tutils.set_random_master_port()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer_options = dict(
default_root_dir=tmpdir,
@@ -62,7 +62,7 @@ def test_amp_gpu_ddp_slurm_managed(tmpdir):
tutils.set_random_master_port()
os.environ['SLURM_LOCALID'] = str(0)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# exp file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -103,7 +103,7 @@ def test_cpu_model_with_amp(tmpdir):
precision=16
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
with pytest.raises((MisconfigurationException, ModuleNotFoundError)):
tutils.run_model_test(trainer_options, model, on_gpu=False)
+12 -12
View File
@@ -1,5 +1,5 @@
from collections import namedtuple
import platform
from collections import namedtuple
import pytest
import torch
@@ -24,7 +24,7 @@ def test_early_stopping_cpu_model(tmpdir):
val_percent_check=0.1,
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
tutils.run_model_test(trainer_options, model, on_gpu=False)
# test freeze on cpu
@@ -53,7 +53,7 @@ def test_multi_cpu_model_ddp(tmpdir):
distributed_backend='ddp_cpu'
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
tutils.run_model_test(trainer_options, model, on_gpu=False)
@@ -68,7 +68,7 @@ def test_lbfgs_cpu_model(tmpdir):
val_percent_check=0.2,
)
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
setattr(hparams, 'optimizer_name', 'lbfgs')
setattr(hparams, 'learning_rate', 0.002)
model = EvalModelTemplate(hparams)
@@ -88,7 +88,7 @@ def test_default_logger_callbacks_cpu_model(tmpdir):
val_percent_check=0.01,
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
tutils.run_model_test_without_loggers(trainer_options, model)
# test freeze on cpu
@@ -98,7 +98,7 @@ def test_default_logger_callbacks_cpu_model(tmpdir):
def test_running_test_after_fitting(tmpdir):
"""Verify test() on fitted model."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -129,7 +129,7 @@ def test_running_test_after_fitting(tmpdir):
def test_running_test_no_val(tmpdir):
"""Verify `test()` works on a model with no `val_loader`."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -207,7 +207,7 @@ def test_single_gpu_batch_parse():
def test_simple_cpu(tmpdir):
"""Verify continue training session on CPU."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# fit model
trainer = Trainer(
@@ -232,7 +232,7 @@ def test_cpu_model(tmpdir):
val_percent_check=0.4
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
tutils.run_model_test(trainer_options, model, on_gpu=False)
@@ -251,7 +251,7 @@ def test_all_features_cpu_model(tmpdir):
val_percent_check=0.4
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
tutils.run_model_test(trainer_options, model, on_gpu=False)
@@ -302,7 +302,7 @@ def test_tbptt_cpu_model(tmpdir):
sampler=None,
)
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
hparams.batch_size = batch_size
hparams.in_features = truncated_bptt_steps
hparams.hidden_dim = truncated_bptt_steps
@@ -336,5 +336,5 @@ def test_single_gpu_model(tmpdir):
gpus=1
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
tutils.run_model_test(trainer_options, model)
+4 -4
View File
@@ -30,7 +30,7 @@ def test_multi_gpu_model(tmpdir, backend):
distributed_backend=backend,
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# tutils.run_model_test(trainer_options, model)
trainer = Trainer(**trainer_options)
result = trainer.fit(model)
@@ -53,7 +53,7 @@ def test_ddp_all_dataloaders_passed_to_fit(tmpdir):
gpus=[0, 1],
distributed_backend='ddp')
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
fit_options = dict(train_dataloader=model.train_dataloader(),
val_dataloaders=model.val_dataloader())
@@ -64,7 +64,7 @@ def test_ddp_all_dataloaders_passed_to_fit(tmpdir):
def test_cpu_slurm_save_load(tmpdir):
"""Verify model save/load/checkpoint on CPU."""
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
# logger file to get meta
@@ -142,7 +142,7 @@ def test_multi_gpu_none_backend(tmpdir):
gpus='-1'
)
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
with pytest.warns(UserWarning):
tutils.run_model_test(trainer_options, model)
+1 -1
View File
@@ -14,7 +14,7 @@ def test_on_before_zero_grad_called(max_steps):
def on_before_zero_grad(self, optimizer):
self.on_before_zero_grad_called += 1
model = CurrentTestModel(tutils.get_default_hparams())
model = CurrentTestModel()
trainer = Trainer(
max_steps=max_steps,
+3 -4
View File
@@ -8,9 +8,8 @@ import sys
import pytest
import torch
from pytorch_lightning import Trainer
import tests.base.utils as tutils
from pytorch_lightning import Trainer
from tests.base import EvalModelTemplate
from tests.base.models import TestGAN
@@ -121,7 +120,7 @@ def test_horovod_transfer_batch_to_gpu(tmpdir):
assert str(y.device) != 'cpu'
return super(TestTrainingStepModel, self).validation_step(batch, *args, **kwargs)
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = TestTrainingStepModel(hparams)
trainer_options = dict(
@@ -139,7 +138,7 @@ def test_horovod_transfer_batch_to_gpu(tmpdir):
@pytest.mark.skipif(sys.version_info >= (3, 8), reason="Horovod not yet supported in Python 3.8")
@pytest.mark.skipif(platform.system() == "Windows", reason="Horovod is not supported on Windows")
def test_horovod_multi_optimizer(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = TestGAN(hparams)
trainer_options = dict(
+1 -1
View File
@@ -27,7 +27,7 @@ def test_training_epoch_end_metrics_collection(tmpdir):
}
}
model = CurrentModel(tutils.get_default_hparams())
model = CurrentModel()
trainer = Trainer(
max_epochs=num_epochs,
default_root_dir=tmpdir,
+7 -9
View File
@@ -19,7 +19,7 @@ def test_running_test_pretrained_model_distrib(tmpdir, backend):
"""Verify `test()` on pretrained model."""
tutils.set_random_master_port()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# exp file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -67,7 +67,7 @@ def test_running_test_pretrained_model_distrib(tmpdir, backend):
def test_running_test_pretrained_model_cpu(tmpdir):
"""Verify test() on pretrained model."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -103,7 +103,7 @@ def test_running_test_pretrained_model_cpu(tmpdir):
def test_load_model_from_checkpoint(tmpdir):
"""Verify test() on pretrained model."""
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
trainer_options = dict(
@@ -145,7 +145,7 @@ def test_load_model_from_checkpoint(tmpdir):
@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="test requires multi-GPU machine")
def test_dp_resume(tmpdir):
"""Make sure DP continues training correctly."""
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
trainer_options = dict(
@@ -217,7 +217,7 @@ def test_dp_resume(tmpdir):
def test_model_saving_loading(tmpdir):
"""Tests use case where trainer saves the model, and user loads it from tags independently."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -284,13 +284,11 @@ def test_load_model_with_missing_hparams(tmpdir):
class CurrentModelWithoutHparams(EvalModelTemplate):
def __init__(self):
hparams = tutils.get_default_hparams()
super().__init__(hparams)
super().__init__()
class CurrentModelUnusedHparams(EvalModelTemplate):
def __init__(self, hparams):
hparams = tutils.get_default_hparams()
super().__init__(hparams)
super().__init__()
model = CurrentModelWithoutHparams()
trainer.fit(model)
+8 -8
View File
@@ -6,7 +6,7 @@ import pytest
from pytorch_lightning import Trainer
import tests.base.utils as tutils
from tests.base import TestModelBase, LightTrainDataloader, LightEmptyTestStep
from tests.base import EvalModelTemplate
def _soft_unimport_module(str_module):
@@ -120,11 +120,11 @@ def test_tbd_remove_in_v0_9_0_module_imports():
from pytorch_lightning.logging.wandb import WandbLogger # noqa: F402
class ModelVer0_6(LightTrainDataloader, LightEmptyTestStep, TestModelBase):
class ModelVer0_6(EvalModelTemplate):
# todo: this shall not be needed while evaluate asks for dataloader explicitly
def val_dataloader(self):
return self._dataloader(train=False)
return self.dataloader(train=False)
def validation_step(self, batch, batch_idx, *args, **kwargs):
return {'val_loss': 0.6}
@@ -133,17 +133,17 @@ class ModelVer0_6(LightTrainDataloader, LightEmptyTestStep, TestModelBase):
return {'val_loss': 0.6}
def test_dataloader(self):
return self._dataloader(train=False)
return self.dataloader(train=False)
def test_end(self, outputs):
return {'test_loss': 0.6}
class ModelVer0_7(LightTrainDataloader, LightEmptyTestStep, TestModelBase):
class ModelVer0_7(EvalModelTemplate):
# todo: this shall not be needed while evaluate asks for dataloader explicitly
def val_dataloader(self):
return self._dataloader(train=False)
return self.dataloader(train=False)
def validation_step(self, batch, batch_idx, *args, **kwargs):
return {'val_loss': 0.7}
@@ -152,14 +152,14 @@ class ModelVer0_7(LightTrainDataloader, LightEmptyTestStep, TestModelBase):
return {'val_loss': 0.7}
def test_dataloader(self):
return self._dataloader(train=False)
return self.dataloader(train=False)
def test_end(self, outputs):
return {'test_loss': 0.7}
def test_tbd_remove_in_v1_0_0_model_hooks():
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = ModelVer0_6(hparams)
+1
View File
@@ -4,6 +4,7 @@ from pathlib import Path
import numpy as np
import pytest
from pytorch_lightning.profiler import AdvancedProfiler, SimpleProfiler
PROFILER_OVERHEAD_MAX_TOLERANCE = 0.0001
+5 -4
View File
@@ -5,6 +5,7 @@ from pytorch_lightning import Trainer
from pytorch_lightning.utilities.exceptions import MisconfigurationException
from tests.base import EvalModelTemplate
# TODO: add matching messages
@@ -14,7 +15,7 @@ def test_wrong_train_setting(tmpdir):
* Test that an error is thrown when no `training_step()` is defined
"""
tutils.reset_seed()
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
trainer = Trainer(default_root_dir=tmpdir, max_epochs=1)
with pytest.raises(MisconfigurationException):
@@ -34,7 +35,7 @@ def test_wrong_configure_optimizers(tmpdir):
trainer = Trainer(default_root_dir=tmpdir, max_epochs=1)
with pytest.raises(MisconfigurationException):
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.configure_optimizers = None
trainer.fit(model)
@@ -47,7 +48,7 @@ def test_wrong_validation_settings(tmpdir):
* error if `validation_step()` is overridden but `val_dataloader()` is not
"""
tutils.reset_seed()
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
trainer = Trainer(default_root_dir=tmpdir, max_epochs=1)
# check val_dataloader -> val_step
@@ -76,7 +77,7 @@ def test_wrong_test_settigs(tmpdir):
throw warning if `test_epoch_end()` is not defined
* error if `test_step()` is overridden but `test_dataloader()` is not
"""
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
trainer = Trainer(default_root_dir=tmpdir, max_epochs=1)
# ----------------
+17 -17
View File
@@ -19,7 +19,7 @@ from tests.base import EvalModelTemplate
])
def test_dataloader_config_errors(tmpdir, dataloader_options):
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# fit model
trainer = Trainer(
@@ -35,7 +35,7 @@ def test_dataloader_config_errors(tmpdir, dataloader_options):
def test_multiple_val_dataloader(tmpdir):
"""Verify multiple val_dataloader."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.val_dataloader = model.val_dataloader__multiple
model.validation_step = model.validation_step__multiple_dataloaders
@@ -63,7 +63,7 @@ def test_multiple_val_dataloader(tmpdir):
def test_multiple_test_dataloader(tmpdir):
"""Verify multiple test_dataloader."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.test_dataloader = model.test_dataloader__multiple
model.test_step = model.test_step__multiple_dataloaders
@@ -93,7 +93,7 @@ def test_train_dataloader_passed_to_fit(tmpdir):
"""Verify that train dataloader can be passed to fit """
# only train passed to fit
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer = Trainer(
default_root_dir=tmpdir,
max_epochs=1,
@@ -110,7 +110,7 @@ def test_train_val_dataloaders_passed_to_fit(tmpdir):
""" Verify that train & val dataloader can be passed to fit """
# train, val passed to fit
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer = Trainer(
default_root_dir=tmpdir,
max_epochs=1,
@@ -129,7 +129,7 @@ def test_train_val_dataloaders_passed_to_fit(tmpdir):
def test_all_dataloaders_passed_to_fit(tmpdir):
"""Verify train, val & test dataloader(s) can be passed to fit and test method"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# train, val and test passed to fit
trainer = Trainer(
@@ -155,7 +155,7 @@ def test_all_dataloaders_passed_to_fit(tmpdir):
def test_multiple_dataloaders_passed_to_fit(tmpdir):
"""Verify that multiple val & test dataloaders can be passed to fit."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.validation_step = model.validation_step__multiple_dataloaders
model.test_step = model.test_step__multiple_dataloaders
@@ -184,7 +184,7 @@ def test_multiple_dataloaders_passed_to_fit(tmpdir):
def test_mixing_of_dataloader_options(tmpdir):
"""Verify that dataloaders can be passed to fit"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer_options = dict(
default_root_dir=tmpdir,
@@ -212,7 +212,7 @@ def test_mixing_of_dataloader_options(tmpdir):
def test_train_inf_dataloader_error(tmpdir):
"""Test inf train data loader (e.g. IterableDataset)"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.train_dataloader = model.train_dataloader__infinite
trainer = Trainer(default_root_dir=tmpdir, max_epochs=1, val_check_interval=0.5)
@@ -223,7 +223,7 @@ def test_train_inf_dataloader_error(tmpdir):
def test_val_inf_dataloader_error(tmpdir):
"""Test inf train data loader (e.g. IterableDataset)"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.val_dataloader = model.val_dataloader__infinite
trainer = Trainer(default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.5)
@@ -234,7 +234,7 @@ def test_val_inf_dataloader_error(tmpdir):
def test_test_inf_dataloader_error(tmpdir):
"""Test inf train data loader (e.g. IterableDataset)"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.test_dataloader = model.test_dataloader__infinite
trainer = Trainer(default_root_dir=tmpdir, max_epochs=1, test_percent_check=0.5)
@@ -247,7 +247,7 @@ def test_test_inf_dataloader_error(tmpdir):
def test_inf_train_dataloader(tmpdir, check_interval):
"""Test inf train data loader (e.g. IterableDataset)"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.train_dataloader = model.train_dataloader__infinite
trainer = Trainer(
@@ -264,7 +264,7 @@ def test_inf_train_dataloader(tmpdir, check_interval):
def test_inf_val_dataloader(tmpdir, check_interval):
"""Test inf val data loader (e.g. IterableDataset)"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.val_dataloader = model.val_dataloader__infinite
# logger file to get meta
@@ -283,7 +283,7 @@ def test_inf_val_dataloader(tmpdir, check_interval):
def test_inf_test_dataloader(tmpdir, check_interval):
"""Test inf test data loader (e.g. IterableDataset)"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.test_dataloader = model.test_dataloader__infinite
# logger file to get meta
@@ -301,7 +301,7 @@ def test_inf_test_dataloader(tmpdir, check_interval):
def test_error_on_zero_len_dataloader(tmpdir):
""" Test that error is raised if a zero-length dataloader is defined """
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.train_dataloader = model.train_dataloader__zero_length
# fit model
@@ -318,7 +318,7 @@ def test_error_on_zero_len_dataloader(tmpdir):
def test_warning_with_few_workers(tmpdir):
""" Test that error is raised if dataloader with only a few workers is used """
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
trainer_options = dict(
@@ -411,7 +411,7 @@ def test_batch_size_smaller_than_num_gpus():
)
return dataloader
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
hparams.batch_size = batch_size
model = CurrentTestModel(hparams)
+6 -6
View File
@@ -10,7 +10,7 @@ from tests.base import EvalModelTemplate
def test_error_on_more_than_1_optimizer(tmpdir):
""" Check that error is thrown when more than 1 optimizer is passed """
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.configure_optimizers = model.configure_optimizers__multiple_schedulers
# logger file to get meta
@@ -26,7 +26,7 @@ def test_error_on_more_than_1_optimizer(tmpdir):
def test_model_reset_correctly(tmpdir):
""" Check that model weights are correctly reset after lr_find() """
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
trainer = Trainer(
@@ -48,7 +48,7 @@ def test_model_reset_correctly(tmpdir):
def test_trainer_reset_correctly(tmpdir):
""" Check that all trainer parameters are reset correctly after lr_find() """
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
trainer = Trainer(
@@ -77,7 +77,7 @@ def test_trainer_reset_correctly(tmpdir):
def test_trainer_arg_bool(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
before_lr = hparams.learning_rate
@@ -96,7 +96,7 @@ def test_trainer_arg_bool(tmpdir):
def test_trainer_arg_str(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
hparams.__dict__['my_fancy_lr'] = 1.0 # update with non-standard field
model = EvalModelTemplate(hparams)
@@ -116,7 +116,7 @@ def test_trainer_arg_str(tmpdir):
def test_call_to_trainer_method(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
before_lr = hparams.learning_rate
+8 -8
View File
@@ -9,7 +9,7 @@ from tests.base import EvalModelTemplate
def test_optimizer_with_scheduling(tmpdir):
""" Verify that learning rate scheduling is working """
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
model.configure_optimizers = model.configure_optimizers__single_scheduler
@@ -40,7 +40,7 @@ def test_optimizer_with_scheduling(tmpdir):
def test_multi_optimizer_with_scheduling(tmpdir):
""" Verify that learning rate scheduling is working """
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
model.configure_optimizers = model.configure_optimizers__multiple_schedulers
@@ -75,7 +75,7 @@ def test_multi_optimizer_with_scheduling(tmpdir):
def test_multi_optimizer_with_scheduling_stepping(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
model.configure_optimizers = model.configure_optimizers__multiple_schedulers
@@ -114,7 +114,7 @@ def test_multi_optimizer_with_scheduling_stepping(tmpdir):
def test_reduce_lr_on_plateau_scheduling(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
model.configure_optimizers = model.configure_optimizers__reduce_lr_on_plateau
@@ -137,7 +137,7 @@ def test_reduce_lr_on_plateau_scheduling(tmpdir):
def test_optimizer_return_options():
trainer = Trainer()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# single optimizer
opt_a = torch.optim.Adam(model.parameters(), lr=0.002)
@@ -195,7 +195,7 @@ def test_none_optimizer_warning():
trainer = Trainer()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.configure_optimizers = lambda: None
with pytest.warns(UserWarning, match='will run with no optimizer'):
@@ -204,7 +204,7 @@ def test_none_optimizer_warning():
def test_none_optimizer(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
model.configure_optimizers = model.configure_optimizers__empty
@@ -231,7 +231,7 @@ def test_configure_optimizer_from_dict(tmpdir):
}
return config
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = CurrentModel(hparams)
# fit model
+16 -16
View File
@@ -20,12 +20,12 @@ from tests.base import EvalModelTemplate
def test_model_pickle(tmpdir):
import pickle
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
pickle.dumps(model)
def test_hparams_save_load(tmpdir):
model = EvalModelTemplate(vars(tutils.get_default_hparams()))
model = EvalModelTemplate(vars(EvalModelTemplate.get_default_hparams()))
trainer = Trainer(
default_root_dir=tmpdir,
@@ -46,7 +46,7 @@ def test_hparams_save_load(tmpdir):
def test_no_val_module(tmpdir):
"""Tests use case where trainer saves the model, and user loads it from tags independently."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -79,7 +79,7 @@ def test_no_val_module(tmpdir):
def test_no_val_end_module(tmpdir):
"""Tests use case where trainer saves the model, and user loads it from tags independently."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
logger = tutils.get_default_logger(tmpdir)
@@ -167,7 +167,7 @@ def test_gradient_accumulation_scheduling(tmpdir):
# clear gradients
optimizer.zero_grad()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
schedule = {1: 2, 3: 4}
trainer = Trainer(accumulate_grad_batches=schedule,
@@ -185,7 +185,7 @@ def test_gradient_accumulation_scheduling(tmpdir):
def test_loading_meta_tags(tmpdir):
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
# save tags
logger = tutils.get_default_logger(tmpdir)
@@ -266,7 +266,7 @@ def test_model_checkpoint_options(tmpdir, save_top_k, file_prefix, expected_file
def test_model_freeze_unfreeze():
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.freeze()
model.unfreeze()
@@ -275,7 +275,7 @@ def test_model_freeze_unfreeze():
def test_resume_from_checkpoint_epoch_restored(tmpdir):
"""Verify resuming from checkpoint runs the right number of epochs"""
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
def _new_model():
# Create a model that tracks epochs and batches seen
@@ -340,7 +340,7 @@ def test_resume_from_checkpoint_epoch_restored(tmpdir):
def _init_steps_model():
"""private method for initializing a model with 5% train epochs"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# define train epoch to 5% of data
train_percent = 0.5
@@ -429,7 +429,7 @@ def test_trainer_min_steps_and_epochs(tmpdir):
def test_benchmark_option(tmpdir):
"""Verify benchmark option."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.val_dataloader = model.val_dataloader__multiple
# verify torch.backends.cudnn.benchmark is not turned on
@@ -452,7 +452,7 @@ def test_benchmark_option(tmpdir):
def test_testpass_overrides(tmpdir):
# todo: check duplicated tests against trainer_checks
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
# Misconfig when neither test_step or test_end is implemented
with pytest.raises(MisconfigurationException, match='.*not implement `test_dataloader`.*'):
@@ -491,7 +491,7 @@ def test_disabled_validation():
self.validation_epoch_end_invoked = True
return super().validation_epoch_end(*args, **kwargs)
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = CurrentModel(hparams)
trainer_options = dict(
@@ -541,7 +541,7 @@ def test_nan_loss_detection(tmpdir):
output /= 0
return output
model = CurrentModel(tutils.get_default_hparams())
model = CurrentModel()
# fit model
trainer = Trainer(
@@ -568,7 +568,7 @@ def test_nan_params_detection(tmpdir):
# simulate parameter that became nan
torch.nn.init.constant_(self.c_d1.bias, math.nan)
model = CurrentModel(tutils.get_default_hparams())
model = CurrentModel()
trainer = Trainer(
default_root_dir=tmpdir,
max_steps=(model.test_batch_nan + 1),
@@ -587,7 +587,7 @@ def test_nan_params_detection(tmpdir):
def test_trainer_interrupted_flag(tmpdir):
"""Test the flag denoting that a user interrupted training."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
class InterruptCallback(Callback):
def __init__(self):
@@ -617,7 +617,7 @@ def test_gradient_clipping(tmpdir):
Test gradient clipping
"""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# test that gradient is clipped correctly
def _optimizer_step(*args, **kwargs):
+1 -1
View File
@@ -1,7 +1,7 @@
import inspect
import pickle
from argparse import ArgumentParser, Namespace
from unittest import mock
import pickle
import pytest
+5 -5
View File
@@ -11,7 +11,7 @@ def test_model_reset_correctly(tmpdir):
""" Check that model weights are correctly reset after scaling batch size. """
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
trainer = Trainer(
@@ -34,7 +34,7 @@ def test_trainer_reset_correctly(tmpdir):
""" Check that all trainer parameters are reset correctly after scaling batch size. """
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
# logger file to get meta
trainer = Trainer(
@@ -71,7 +71,7 @@ def test_trainer_arg(tmpdir, scale_arg):
""" Check that trainer arg works with bool input. """
tutils.reset_seed()
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
before_batch_size = hparams.batch_size
@@ -93,7 +93,7 @@ def test_call_to_trainer_method(tmpdir, scale_method):
""" Test that calling the trainer method itself works. """
tutils.reset_seed()
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
before_batch_size = hparams.batch_size
@@ -116,7 +116,7 @@ def test_error_on_dataloader_passed_to_fit(tmpdir):
if a train dataloader is passed to fit """
# only train passed to fit
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer = Trainer(
default_root_dir=tmpdir,
max_epochs=1,