mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Tests: refactor cleanup (#1744)
* wip * cleaning * optim imports * - * default hparams * fix restore * fix imports
This commit is contained in:
+2
-59
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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):
|
||||
@@ -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
@@ -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])
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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):
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,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)
|
||||
|
||||
# ----------------
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,7 +1,7 @@
|
||||
import inspect
|
||||
import pickle
|
||||
from argparse import ArgumentParser, Namespace
|
||||
from unittest import mock
|
||||
import pickle
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user