diff --git a/pl_examples/basic_examples/cpu_template.py b/pl_examples/basic_examples/cpu_template.py index 54e16208..fd750bc2 100644 --- a/pl_examples/basic_examples/cpu_template.py +++ b/pl_examples/basic_examples/cpu_template.py @@ -2,12 +2,13 @@ Runs a model on a single node across N-gpus. """ import os +from argparse import ArgumentParser + import numpy as np import torch -from argparse import ArgumentParser -from pytorch_lightning import Trainer from pl_examples.basic_examples.lightning_module_template import LightningTemplateModel +from pytorch_lightning import Trainer SEED = 2334 torch.manual_seed(SEED) diff --git a/pl_examples/basic_examples/gpu_template.py b/pl_examples/basic_examples/gpu_template.py index 9b15f463..d95b9ab8 100644 --- a/pl_examples/basic_examples/gpu_template.py +++ b/pl_examples/basic_examples/gpu_template.py @@ -2,12 +2,13 @@ Runs a model on a single node across N-gpus. """ import os +from argparse import ArgumentParser + import numpy as np import torch -from argparse import ArgumentParser -from pytorch_lightning import Trainer from pl_examples.basic_examples.lightning_module_template import LightningTemplateModel +from pytorch_lightning import Trainer SEED = 2334 torch.manual_seed(SEED) diff --git a/pl_examples/basic_examples/lightning_module_template.py b/pl_examples/basic_examples/lightning_module_template.py index 517bf9a0..f773fc81 100644 --- a/pl_examples/basic_examples/lightning_module_template.py +++ b/pl_examples/basic_examples/lightning_module_template.py @@ -2,16 +2,17 @@ Example template for defining a system """ import os -from collections import OrderedDict -import torch.nn as nn -from torchvision.datasets import MNIST -import torchvision.transforms as transforms -import torch -import torch.nn.functional as F from argparse import ArgumentParser +from collections import OrderedDict + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torchvision.transforms as transforms from torch import optim from torch.utils.data import DataLoader from torch.utils.data.distributed import DistributedSampler +from torchvision.datasets import MNIST import pytorch_lightning as pl from pytorch_lightning.root_module.root_module import LightningModule diff --git a/pl_examples/domain_templates/gan.py b/pl_examples/domain_templates/gan.py index 29f66768..a8468c4d 100644 --- a/pl_examples/domain_templates/gan.py +++ b/pl_examples/domain_templates/gan.py @@ -6,20 +6,18 @@ After a few epochs, launch tensorboard to see the images being generated at ever tensorboard --logdir default """ -from argparse import ArgumentParser import os -import numpy as np +from argparse import ArgumentParser from collections import OrderedDict -import torchvision -import torchvision.transforms as transforms -from torchvision.datasets import MNIST - -from torch.utils.data import DataLoader - +import numpy as np +import torch import torch.nn as nn import torch.nn.functional as F -import torch +import torchvision +import torchvision.transforms as transforms +from torch.utils.data import DataLoader +from torchvision.datasets import MNIST import pytorch_lightning as pl @@ -178,7 +176,6 @@ class GAN(pl.LightningModule): def main(hparams): - # ------------------------ # 1 INIT LIGHTNING MODEL # ------------------------ diff --git a/pl_examples/multi_node_examples/multi_node_ddp2_demo.py b/pl_examples/multi_node_examples/multi_node_ddp2_demo.py index 1267fbcd..0a4a45db 100644 --- a/pl_examples/multi_node_examples/multi_node_ddp2_demo.py +++ b/pl_examples/multi_node_examples/multi_node_ddp2_demo.py @@ -2,12 +2,13 @@ Multi-node example (GPU) """ import os +from argparse import ArgumentParser + import numpy as np import torch -from argparse import ArgumentParser -from pytorch_lightning import Trainer from pl_examples.basic_examples.lightning_module_template import LightningTemplateModel +from pytorch_lightning import Trainer SEED = 2334 torch.manual_seed(SEED) @@ -41,7 +42,6 @@ def main(hparams): if __name__ == '__main__': - root_dir = os.path.dirname(os.path.realpath(__file__)) parent_parser = ArgumentParser(add_help=False) diff --git a/pl_examples/multi_node_examples/multi_node_ddp_demo.py b/pl_examples/multi_node_examples/multi_node_ddp_demo.py index 7b41ad4c..46877d39 100644 --- a/pl_examples/multi_node_examples/multi_node_ddp_demo.py +++ b/pl_examples/multi_node_examples/multi_node_ddp_demo.py @@ -2,12 +2,13 @@ Multi-node example (GPU) """ import os +from argparse import ArgumentParser + import numpy as np import torch -from argparse import ArgumentParser -from pytorch_lightning import Trainer from pl_examples.basic_examples.lightning_module_template import LightningTemplateModel +from pytorch_lightning import Trainer SEED = 2334 torch.manual_seed(SEED) @@ -41,7 +42,6 @@ def main(hparams): if __name__ == '__main__': - root_dir = os.path.dirname(os.path.realpath(__file__)) parent_parser = ArgumentParser(add_help=False) diff --git a/pytorch_lightning/__init__.py b/pytorch_lightning/__init__.py index 7bdaf7b5..d65aff02 100644 --- a/pytorch_lightning/__init__.py +++ b/pytorch_lightning/__init__.py @@ -1,6 +1,6 @@ -from .trainer.trainer import Trainer -from .root_module.root_module import LightningModule from .root_module.decorators import data_loader +from .root_module.root_module import LightningModule +from .trainer.trainer import Trainer __all__ = [ 'Trainer', diff --git a/pytorch_lightning/callbacks/pt_callbacks.py b/pytorch_lightning/callbacks/pt_callbacks.py index 12d8919c..d5d80d24 100644 --- a/pytorch_lightning/callbacks/pt_callbacks.py +++ b/pytorch_lightning/callbacks/pt_callbacks.py @@ -1,6 +1,5 @@ import os import shutil -import warnings import numpy as np @@ -260,6 +259,7 @@ class GradientAccumulationScheduler(Callback): # Arguments scheduling: dict, scheduling in format {epoch: accumulation_factor} """ + def __init__(self, scheduling: dict): if scheduling == {}: # empty dict error raise TypeError("Empty dict cannot be interpreted correct") diff --git a/pytorch_lightning/logging/comet_logger.py b/pytorch_lightning/logging/comet_logger.py index 386001b4..e2ac0c24 100644 --- a/pytorch_lightning/logging/comet_logger.py +++ b/pytorch_lightning/logging/comet_logger.py @@ -1,7 +1,7 @@ -from time import time -from logging import getLogger from os import environ + from comet_ml import Experiment as CometExperiment + from .base import LightningLoggerBase, rank_zero_only # needed to prevent ImportError and duplicated logs. diff --git a/pytorch_lightning/logging/mlflow_logger.py b/pytorch_lightning/logging/mlflow_logger.py index 667488ce..c38b7a16 100644 --- a/pytorch_lightning/logging/mlflow_logger.py +++ b/pytorch_lightning/logging/mlflow_logger.py @@ -1,5 +1,5 @@ -from time import time from logging import getLogger +from time import time import mlflow diff --git a/pytorch_lightning/logging/test_tube_logger.py b/pytorch_lightning/logging/test_tube_logger.py index cd721512..e8e9607f 100644 --- a/pytorch_lightning/logging/test_tube_logger.py +++ b/pytorch_lightning/logging/test_tube_logger.py @@ -1,16 +1,13 @@ -import os.path -from copy import copy +from test_tube import Experiment from .base import LightningLoggerBase, rank_zero_only -from test_tube import Experiment - class TestTubeLogger(LightningLoggerBase): __test__ = False def __init__( - self, save_dir, name="default", description=None, debug=False, + self, save_dir, name="default", description=None, debug=False, version=None, create_git_tag=False ): super().__init__() diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/pt_overrides/override_data_parallel.py index dd3fe2d2..adf3f631 100644 --- a/pytorch_lightning/pt_overrides/override_data_parallel.py +++ b/pytorch_lightning/pt_overrides/override_data_parallel.py @@ -1,11 +1,11 @@ -from torch.nn import DataParallel -from torch.nn.parallel import DistributedDataParallel import itertools +import threading from itertools import chain -import threading import torch from torch.cuda._utils import _get_device_index +from torch.nn import DataParallel +from torch.nn.parallel import DistributedDataParallel def _find_tensors(obj): # pragma: no cover diff --git a/pytorch_lightning/root_module/decorators.py b/pytorch_lightning/root_module/decorators.py index 65062f73..0a87e00f 100644 --- a/pytorch_lightning/root_module/decorators.py +++ b/pytorch_lightning/root_module/decorators.py @@ -17,9 +17,9 @@ def data_loader(fn): try: value = fn(self) # Lazy evaluation, done only once. if ( - value is not None and - not isinstance(value, list) and - fn.__name__ in['test_dataloader', 'val_dataloader'] + value is not None and + not isinstance(value, list) and + fn.__name__ in ['test_dataloader', 'val_dataloader'] ): value = [value] except AttributeError as e: diff --git a/pytorch_lightning/root_module/memory.py b/pytorch_lightning/root_module/memory.py index a4e94e3f..d16d2659 100644 --- a/pytorch_lightning/root_module/memory.py +++ b/pytorch_lightning/root_module/memory.py @@ -3,11 +3,11 @@ Generates a summary of a model's layers and dimensionality ''' import gc - -import torch import subprocess + import numpy as np import pandas as pd +import torch class ModelSummary(object): @@ -146,7 +146,6 @@ class ModelSummary(object): df['Params'] = df['Params'].map(get_human_readable_count) if self.model.example_input_array is not None: - df['In_sizes'] = self.in_sizes df['Out_sizes'] = self.out_sizes diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index 7da13acd..a4f4fb60 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -1,13 +1,13 @@ +import warnings + import torch -from pytorch_lightning.root_module.memory import ModelSummary -from pytorch_lightning.root_module.grads import GradInformation -from pytorch_lightning.trainer.trainer_io import load_hparams_from_tags_csv -from pytorch_lightning.root_module.model_saving import ModelIO -from pytorch_lightning.root_module.hooks import ModelHooks from pytorch_lightning.root_module.decorators import data_loader - -import warnings +from pytorch_lightning.root_module.grads import GradInformation +from pytorch_lightning.root_module.hooks import ModelHooks +from pytorch_lightning.root_module.memory import ModelSummary +from pytorch_lightning.root_module.model_saving import ModelIO +from pytorch_lightning.trainer.trainer_io import load_hparams_from_tags_csv class LightningModule(GradInformation, ModelIO, ModelHooks): diff --git a/pytorch_lightning/testing/lm_test_module.py b/pytorch_lightning/testing/lm_test_module.py index e8901163..dd0992e0 100644 --- a/pytorch_lightning/testing/lm_test_module.py +++ b/pytorch_lightning/testing/lm_test_module.py @@ -1,18 +1,4 @@ -import os -from collections import OrderedDict - import torch -import torch.nn as nn -import torch.nn.functional as F -from torch import optim -from torch.utils.data import DataLoader -from torch.utils.data.distributed import DistributedSampler -from torchvision.datasets import MNIST -from torchvision import transforms -from test_tube import HyperOptArgumentParser - -from pytorch_lightning.root_module.root_module import LightningModule -from pytorch_lightning import data_loader from .lm_test_module_base import LightningTestModelBase from .lm_test_module_mixins import LightningValidationMixin, LightningTestMixin diff --git a/pytorch_lightning/testing/lm_test_module_base.py b/pytorch_lightning/testing/lm_test_module_base.py index 5bd7c0bc..9a6e0ee6 100644 --- a/pytorch_lightning/testing/lm_test_module_base.py +++ b/pytorch_lightning/testing/lm_test_module_base.py @@ -4,15 +4,15 @@ from collections import OrderedDict import torch import torch.nn as nn import torch.nn.functional as F +from test_tube import HyperOptArgumentParser from torch import optim from torch.utils.data import DataLoader from torch.utils.data.distributed import DistributedSampler -from torchvision.datasets import MNIST from torchvision import transforms -from test_tube import HyperOptArgumentParser +from torchvision.datasets import MNIST -from pytorch_lightning.root_module.root_module import LightningModule from pytorch_lightning import data_loader +from pytorch_lightning.root_module.root_module import LightningModule class LightningTestModelBase(LightningModule): diff --git a/pytorch_lightning/testing/lm_test_module_mixins.py b/pytorch_lightning/testing/lm_test_module_mixins.py index 562126db..3815e84e 100644 --- a/pytorch_lightning/testing/lm_test_module_mixins.py +++ b/pytorch_lightning/testing/lm_test_module_mixins.py @@ -1,17 +1,7 @@ -import os from collections import OrderedDict import torch -import torch.nn as nn -import torch.nn.functional as F -from torch import optim -from torch.utils.data import DataLoader -from torch.utils.data.distributed import DistributedSampler -from torchvision.datasets import MNIST -from torchvision import transforms -from test_tube import HyperOptArgumentParser -from pytorch_lightning.root_module.root_module import LightningModule from pytorch_lightning import data_loader diff --git a/pytorch_lightning/trainer/amp_mixin.py b/pytorch_lightning/trainer/amp_mixin.py index 85f2248d..d0790242 100644 --- a/pytorch_lightning/trainer/amp_mixin.py +++ b/pytorch_lightning/trainer/amp_mixin.py @@ -1,5 +1,6 @@ try: from apex import amp + APEX_AVAILABLE = True except ImportError: APEX_AVAILABLE = False diff --git a/pytorch_lightning/trainer/data_loading_mixin.py b/pytorch_lightning/trainer/data_loading_mixin.py index ce511b56..066c21dc 100644 --- a/pytorch_lightning/trainer/data_loading_mixin.py +++ b/pytorch_lightning/trainer/data_loading_mixin.py @@ -1,13 +1,14 @@ import warnings -from torch.utils.data.distributed import DistributedSampler import torch.distributed as dist from torch.utils.data import IterableDataset +from torch.utils.data.distributed import DistributedSampler from pytorch_lightning.utilities.debugging import MisconfigurationException try: from apex import amp + APEX_AVAILABLE = True except ImportError: APEX_AVAILABLE = False diff --git a/pytorch_lightning/trainer/ddp_mixin.py b/pytorch_lightning/trainer/ddp_mixin.py index 1b9cac56..36345b74 100644 --- a/pytorch_lightning/trainer/ddp_mixin.py +++ b/pytorch_lightning/trainer/ddp_mixin.py @@ -4,12 +4,13 @@ import warnings import torch import torch.distributed as dist + from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel from pytorch_lightning.utilities.debugging import MisconfigurationException - try: from apex import amp + APEX_AVAILABLE = True except ImportError: APEX_AVAILABLE = False diff --git a/pytorch_lightning/trainer/dp_mixin.py b/pytorch_lightning/trainer/dp_mixin.py index 76403164..96aa5eb3 100644 --- a/pytorch_lightning/trainer/dp_mixin.py +++ b/pytorch_lightning/trainer/dp_mixin.py @@ -1,17 +1,12 @@ -import os -import re -import signal -import pdb -from subprocess import call - import torch -import torch.distributed as dist + from pytorch_lightning.pt_overrides.override_data_parallel import ( LightningDistributedDataParallel, LightningDataParallel) from pytorch_lightning.utilities.debugging import MisconfigurationException try: from apex import amp + APEX_AVAILABLE = True except ImportError: APEX_AVAILABLE = False diff --git a/pytorch_lightning/trainer/evaluation_loop_mixin.py b/pytorch_lightning/trainer/evaluation_loop_mixin.py index 0373c188..c3cc51ef 100644 --- a/pytorch_lightning/trainer/evaluation_loop_mixin.py +++ b/pytorch_lightning/trainer/evaluation_loop_mixin.py @@ -1,4 +1,5 @@ import torch + from pytorch_lightning.utilities.debugging import MisconfigurationException diff --git a/pytorch_lightning/trainer/logging_mixin.py b/pytorch_lightning/trainer/logging_mixin.py index a66541bf..f437ee99 100644 --- a/pytorch_lightning/trainer/logging_mixin.py +++ b/pytorch_lightning/trainer/logging_mixin.py @@ -1,4 +1,5 @@ import torch + from pytorch_lightning.root_module import memory diff --git a/pytorch_lightning/trainer/train_loop_mixin.py b/pytorch_lightning/trainer/train_loop_mixin.py index 5e54d8d6..e51a43d9 100644 --- a/pytorch_lightning/trainer/train_loop_mixin.py +++ b/pytorch_lightning/trainer/train_loop_mixin.py @@ -2,6 +2,7 @@ import numpy as np try: from apex import amp + APEX_AVAILABLE = True except ImportError: APEX_AVAILABLE = False @@ -112,7 +113,6 @@ class TrainerTrainLoopMixin(object): # when metrics should be logged should_log_metrics = batch_nb % self.row_log_interval == 0 or early_stop_epoch if should_log_metrics or self.fast_dev_run: - # logs user requested information to logger self.log_metrics(batch_step_metrics, grad_norm_dic) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index de71ef17..a9ae29b7 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -5,29 +5,28 @@ The trainer handles all the logic for running a val loop, training loop, distrib import os import warnings -import tqdm import torch -import torch.multiprocessing as mp import torch.distributed as dist +import torch.multiprocessing as mp +import tqdm from torch.optim.optimizer import Optimizer -from pytorch_lightning.trainer.trainer_io import TrainerIOMixin +from pytorch_lightning.trainer.amp_mixin import TrainerAMPMixin +from pytorch_lightning.trainer.callback_config_mixin import TrainerCallbackConfigMixin +from pytorch_lightning.trainer.data_loading_mixin import TrainerDataLoadingMixin from pytorch_lightning.trainer.ddp_mixin import TrainerDDPMixin from pytorch_lightning.trainer.dp_mixin import TrainerDPMixin -from pytorch_lightning.trainer.amp_mixin import TrainerAMPMixin -from pytorch_lightning.trainer.data_loading_mixin import TrainerDataLoadingMixin from pytorch_lightning.trainer.evaluation_loop_mixin import TrainerEvaluationLoopMixin -from pytorch_lightning.trainer.train_loop_mixin import TrainerTrainLoopMixin from pytorch_lightning.trainer.logging_mixin import TrainerLoggingMixin -from pytorch_lightning.trainer.training_tricks_mixin import TrainerTrainingTricksMixin -from pytorch_lightning.trainer.callback_config_mixin import TrainerCallbackConfigMixin from pytorch_lightning.trainer.model_hooks_mixin import TrainerModelHooksMixin - +from pytorch_lightning.trainer.train_loop_mixin import TrainerTrainLoopMixin +from pytorch_lightning.trainer.trainer_io import TrainerIOMixin +from pytorch_lightning.trainer.training_tricks_mixin import TrainerTrainingTricksMixin from pytorch_lightning.utilities.debugging import MisconfigurationException -import pdb try: from apex import amp + APEX_AVAILABLE = True except ImportError: APEX_AVAILABLE = False @@ -332,7 +331,7 @@ class Trainer(TrainerIOMixin, task = int(os.environ['SLURM_LOCALID']) self.ddp_train(task, model) else: - mp.spawn(self.ddp_train, nprocs=self.num_gpus, args=(model, )) + mp.spawn(self.ddp_train, nprocs=self.num_gpus, args=(model,)) # 1 gpu or dp option triggers training using DP module # easier to avoid NCCL issues diff --git a/pytorch_lightning/trainer/trainer_io.py b/pytorch_lightning/trainer/trainer_io.py index 65934837..d1c99767 100644 --- a/pytorch_lightning/trainer/trainer_io.py +++ b/pytorch_lightning/trainer/trainer_io.py @@ -1,11 +1,11 @@ import os import re import signal -import pdb from subprocess import call import torch import torch.distributed as dist + from pytorch_lightning.pt_overrides.override_data_parallel import ( LightningDistributedDataParallel, LightningDataParallel) diff --git a/pytorch_lightning/utilities/arg_parse.py b/pytorch_lightning/utilities/arg_parse.py index fbf3ec2f..73843e52 100644 --- a/pytorch_lightning/utilities/arg_parse.py +++ b/pytorch_lightning/utilities/arg_parse.py @@ -7,7 +7,6 @@ import os def add_default_args(parser, root_dir, rand_seed=None, possible_model_names=None): - # training, test, val check intervals parser.add_argument('--eval_test_set', dest='eval_test_set', action='store_true', help='true = run test set also') diff --git a/tests/debug.py b/tests/debug.py index 3a8ccb6c..c8adf8c0 100644 --- a/tests/debug.py +++ b/tests/debug.py @@ -1,19 +1,11 @@ -from pytorch_lightning import Trainer -from pl_examples import LightningTemplateModel -from pytorch_lightning.testing import LightningTestModel -from argparse import Namespace -from test_tube import Experiment -from pytorch_lightning.callbacks import ModelCheckpoint -import os -import shutil - -import pytorch_lightning as pl import torch from torch.nn import functional as F from torch.utils.data import DataLoader from torchvision.datasets import MNIST -import numpy as np -import pdb + +import pytorch_lightning as pl + + # from test_models import assert_ok_test_acc, load_model, \ # clear_save_dir, get_test_tube_logger, get_hparams, init_save_dir, \ # init_checkpoint_callback, reset_seed, set_random_master_port diff --git a/tests/test_logging.py b/tests/test_logging.py index cc3489e8..9db65750 100644 --- a/tests/test_logging.py +++ b/tests/test_logging.py @@ -1,16 +1,12 @@ -import os.path import pickle -import shutil import numpy as np import torch + from pytorch_lightning import Trainer - from pytorch_lightning.testing import LightningTestModel -from pytorch_lightning.logging import LightningLoggerBase, rank_zero_only from .test_models import get_hparams, get_test_tube_logger, init_save_dir, clear_save_dir - RANDOM_FILE_PATHS = list(np.random.randint(12000, 19000, 1000)) ROOT_SEED = 1234 torch.manual_seed(ROOT_SEED) diff --git a/tests/test_models.py b/tests/test_models.py index a500f224..ef99b968 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -3,34 +3,30 @@ import shutil import warnings from argparse import Namespace -import pytest import numpy as np +import pytest import torch +from pl_examples import LightningTemplateModel # sys.path += [os.path.abspath('..'), os.path.abspath('../..')] from pytorch_lightning import Trainer +from pytorch_lightning.callbacks import ( + ModelCheckpoint, + EarlyStopping, +) +from pytorch_lightning.logging import TestTubeLogger +from pytorch_lightning.root_module import memory from pytorch_lightning.testing import ( LightningTestModel, LightningTestModelBase, - LightningValidationMixin, LightningValidationStepMixin, LightningValidationMultipleDataloadersMixin, LightningTestMixin, LightningTestMultipleDataloadersMixin, ) -from pytorch_lightning.callbacks import ( - ModelCheckpoint, - EarlyStopping, - GradientAccumulationScheduler, -) -from pytorch_lightning.utilities.debugging import MisconfigurationException -from pytorch_lightning.root_module import memory -from pytorch_lightning.root_module import model_saving from pytorch_lightning.trainer import trainer_io -from pytorch_lightning.logging import TestTubeLogger -from pl_examples import LightningTemplateModel from pytorch_lightning.trainer.logging_mixin import TrainerLoggingMixin - +from pytorch_lightning.utilities.debugging import MisconfigurationException # generate a list of random seeds for each test RANDOM_FILE_PATHS = list(np.random.randint(12000, 19000, 1000)) @@ -325,8 +321,10 @@ def test_running_test_without_val(): reset_seed() """Verify test() works on a model with no val_loader""" + class CurrentTestModel(LightningTestMixin, LightningTestModelBase): pass + hparams = get_hparams() model = CurrentTestModel(hparams) @@ -634,6 +632,7 @@ def test_no_val_module(): class CurrentTestModel(LightningTestModelBase): pass + model = CurrentTestModel(hparams) save_dir = init_save_dir() @@ -678,6 +677,7 @@ def test_no_val_end_module(): class CurrentTestModel(LightningValidationStepMixin, LightningTestModelBase): pass + hparams = get_hparams() model = CurrentTestModel(hparams) @@ -1365,6 +1365,7 @@ def test_multiple_val_dataloader(): LightningTestModelBase ): pass + hparams = get_hparams() model = CurrentTestModel(hparams) @@ -1402,6 +1403,7 @@ def test_multiple_test_dataloader(): LightningTestModelBase ): pass + hparams = get_hparams() model = CurrentTestModel(hparams) diff --git a/tox.ini b/tox.ini index 79fa0819..b3a62978 100644 --- a/tox.ini +++ b/tox.ini @@ -21,7 +21,7 @@ log_cli_format = %(message)s log_file = pytest.log log_file_level = DEBUG log_file_format = %(asctime)s [%(levelname)8s] %(message)s (%(filename)s:%(lineno)s) -log_file_date_format=%Y-%m-%d %H:%M:%S +log_file_date_format = %Y-%m-%d %H:%M:%S [testenv] basepython =