From 34bc1493596697a2dfc8c76036921b2bb2fb5013 Mon Sep 17 00:00:00 2001 From: Jirka Borovec Date: Fri, 1 May 2020 16:43:58 +0200 Subject: [PATCH] move unnecessary dict trainer_options (#1469) * move unnecessary dict trainer_options * fix tests * fix tests * formatting * missing --- tests/base/utils.py | 2 - tests/callbacks/test_callbacks.py | 21 ++-- tests/loggers/test_base.py | 15 +-- tests/loggers/test_neptune.py | 4 +- tests/loggers/test_trains.py | 7 +- tests/loggers/test_wandb.py | 8 +- tests/models/test_amp.py | 24 ++--- tests/models/test_cpu.py | 37 +++---- tests/models/test_gpu.py | 33 +++--- tests/models/test_restore.py | 40 ++++---- tests/trainer/test_checks.py | 12 +-- tests/trainer/test_dataloaders.py | 161 +++++++++--------------------- tests/trainer/test_optimizers.py | 40 +++----- tests/trainer/test_trainer.py | 58 ++++------- 14 files changed, 161 insertions(+), 301 deletions(-) diff --git a/tests/base/utils.py b/tests/base/utils.py index fc10d75b..42e6d17d 100644 --- a/tests/base/utils.py +++ b/tests/base/utils.py @@ -27,8 +27,6 @@ def assert_speed_parity(pl_times, pt_times, num_epochs): def run_model_test_without_loggers(trainer_options, model, min_acc=0.50): - # save_dir = trainer_options['default_root_dir'] - # fit model trainer = Trainer(**trainer_options) result = trainer.fit(model) diff --git a/tests/callbacks/test_callbacks.py b/tests/callbacks/test_callbacks.py index a082c5ec..fcd0836f 100644 --- a/tests/callbacks/test_callbacks.py +++ b/tests/callbacks/test_callbacks.py @@ -126,13 +126,13 @@ def test_trainer_callback_system(tmpdir): test_callback = TestCallback() - trainer_options = { - 'callbacks': [test_callback], - 'max_epochs': 1, - 'val_percent_check': 0.1, - 'train_percent_check': 0.2, - 'progress_bar_refresh_rate': 0 - } + trainer_options = dict( + callbacks=[test_callback], + max_epochs=1, + val_percent_check=0.1, + train_percent_check=0.2, + progress_bar_refresh_rate=0, + ) assert not test_callback.on_init_start_called assert not test_callback.on_init_end_called @@ -198,7 +198,7 @@ def test_trainer_callback_system(tmpdir): assert not test_callback.on_test_end_called test_callback = TestCallback() - trainer_options['callbacks'] = [test_callback] + trainer_options.update(callbacks=[test_callback]) trainer = Trainer(**trainer_options) trainer.test(model) @@ -228,14 +228,13 @@ def test_early_stopping_no_val_step(tmpdir): model = ModelWithoutValStep(hparams) stopping = EarlyStopping(monitor='my_train_metric', min_delta=0.1) - trainer_options = dict( + + trainer = Trainer( default_root_dir=tmpdir, early_stop_callback=stopping, overfit_pct=0.20, max_epochs=5, ) - - trainer = Trainer(**trainer_options) result = trainer.fit(model) assert result == 1, 'training failed to complete' diff --git a/tests/loggers/test_base.py b/tests/loggers/test_base.py index 9bcbc2fb..56f6b97b 100644 --- a/tests/loggers/test_base.py +++ b/tests/loggers/test_base.py @@ -65,14 +65,12 @@ def test_custom_logger(tmpdir): logger = CustomLogger() - trainer_options = dict( + trainer = Trainer( max_epochs=1, train_percent_check=0.05, logger=logger, default_root_dir=tmpdir ) - - trainer = Trainer(**trainer_options) result = trainer.fit(model) assert result == 1, "Training failed" assert logger.hparams_logged == hparams @@ -87,14 +85,12 @@ def test_multiple_loggers(tmpdir): logger1 = CustomLogger() logger2 = CustomLogger() - trainer_options = dict( + trainer = Trainer( max_epochs=1, train_percent_check=0.05, logger=[logger1, logger2], default_root_dir=tmpdir ) - - trainer = Trainer(**trainer_options) result = trainer.fit(model) assert result == 1, "Training failed" @@ -113,9 +109,7 @@ def test_multiple_loggers_pickle(tmpdir): logger1 = CustomLogger() logger2 = CustomLogger() - trainer_options = dict(max_epochs=1, logger=[logger1, logger2]) - - trainer = Trainer(**trainer_options) + trainer = Trainer(max_epochs=1, logger=[logger1, logger2]) pkl_bytes = pickle.dumps(trainer) trainer2 = pickle.loads(pkl_bytes) trainer2.logger.log_metrics({"acc": 1.0}, 0) @@ -148,14 +142,13 @@ def test_adding_step_key(tmpdir): model, hparams = tutils.get_default_model() model.validation_epoch_end = _validation_epoch_end model.training_epoch_end = _training_epoch_end - trainer_options = dict( + trainer = Trainer( max_epochs=4, default_root_dir=tmpdir, train_percent_check=0.001, val_percent_check=0.01, num_sanity_val_steps=0, ) - trainer = Trainer(**trainer_options) trainer.logger.log_metrics = _log_metrics_decorator( trainer.logger.log_metrics) trainer.fit(model) diff --git a/tests/loggers/test_neptune.py b/tests/loggers/test_neptune.py index 09f531ab..4cfdd673 100644 --- a/tests/loggers/test_neptune.py +++ b/tests/loggers/test_neptune.py @@ -68,14 +68,12 @@ def test_neptune_leave_open_experiment_after_fit(tmpdir): def _run_training(logger): logger._experiment = MagicMock() - - trainer_options = dict( + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, train_percent_check=0.05, logger=logger ) - trainer = Trainer(**trainer_options) trainer.fit(model) return logger diff --git a/tests/loggers/test_trains.py b/tests/loggers/test_trains.py index 3dafb570..e4ee78c6 100644 --- a/tests/loggers/test_trains.py +++ b/tests/loggers/test_trains.py @@ -18,13 +18,12 @@ def test_trains_logger(tmpdir): web_host='http://integration.trains.allegro.ai:8080', ) logger = TrainsLogger(project_name="lightning_log", task_name="pytorch lightning test") - trainer_options = dict( + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, train_percent_check=0.05, logger=logger ) - trainer = Trainer(**trainer_options) result = trainer.fit(model) print('result finished') @@ -44,13 +43,11 @@ def test_trains_pickle(tmpdir): web_host='http://integration.trains.allegro.ai:8080', ) logger = TrainsLogger(project_name="lightning_log", task_name="pytorch lightning test") - trainer_options = dict( + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, logger=logger ) - - trainer = Trainer(**trainer_options) pkl_bytes = pickle.dumps(trainer) trainer2 = pickle.loads(pkl_bytes) trainer2.logger.log_metrics({"acc": 1.0}) diff --git a/tests/loggers/test_wandb.py b/tests/loggers/test_wandb.py index d2ef6318..bb2739f9 100644 --- a/tests/loggers/test_wandb.py +++ b/tests/loggers/test_wandb.py @@ -47,17 +47,15 @@ def test_wandb_pickle(wandb): logger = WandbLogger(id='the_id', offline=True) - trainer_options = dict(max_epochs=1, logger=logger) - - trainer = Trainer(**trainer_options) + trainer = Trainer(max_epochs=1, logger=logger) # Access the experiment to ensure it's created - trainer.logger.experiment + assert trainer.logger.experiment, 'missing experiment' pkl_bytes = pickle.dumps(trainer) trainer2 = pickle.loads(pkl_bytes) assert os.environ['WANDB_MODE'] == 'dryrun' assert trainer2.logger.__class__.__name__ == WandbLogger.__name__ - _ = trainer2.logger.experiment + assert trainer2.logger.experiment, 'missing experiment' wandb.init.assert_called() assert 'id' in wandb.init.call_args[1] diff --git a/tests/models/test_amp.py b/tests/models/test_amp.py index 9b21e711..81a2325c 100644 --- a/tests/models/test_amp.py +++ b/tests/models/test_amp.py @@ -20,7 +20,7 @@ def test_amp_single_gpu(tmpdir, backend): model, hparams = tutils.get_default_model() - trainer_options = dict( + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, gpus=1, @@ -29,8 +29,6 @@ def test_amp_single_gpu(tmpdir, backend): ) # tutils.run_model_test(trainer_options, model) - - trainer = Trainer(**trainer_options) result = trainer.fit(model) assert result == 1 @@ -74,25 +72,21 @@ def test_amp_gpu_ddp_slurm_managed(tmpdir): hparams = tutils.get_default_hparams() model = LightningTestModel(hparams) - trainer_options = dict( - max_epochs=1, - gpus=[0], - distributed_backend='ddp', - precision=16 - ) - # exp file to get meta logger = tutils.get_default_logger(tmpdir) # exp file to get weights checkpoint = tutils.init_checkpoint_callback(logger) - # add these to the trainer options - trainer_options['checkpoint_callback'] = checkpoint - trainer_options['logger'] = logger - # fit model - trainer = Trainer(**trainer_options) + trainer = Trainer( + max_epochs=1, + gpus=[0], + distributed_backend='ddp', + precision=16, + checkpoint_callback=checkpoint, + logger=logger, + ) trainer.is_slurm_managing_tasks = True result = trainer.fit(model) diff --git a/tests/models/test_cpu.py b/tests/models/test_cpu.py index 0c4ca6e4..e7b422dc 100644 --- a/tests/models/test_cpu.py +++ b/tests/models/test_cpu.py @@ -120,7 +120,8 @@ def test_running_test_after_fitting(tmpdir): # logger file to get weights checkpoint = tutils.init_checkpoint_callback(logger) - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, progress_bar_refresh_rate=0, max_epochs=8, @@ -130,9 +131,6 @@ def test_running_test_after_fitting(tmpdir): checkpoint_callback=checkpoint, logger=logger ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) assert result == 1, 'training failed to complete' @@ -159,7 +157,8 @@ def test_running_test_no_val(tmpdir): # logger file to get weights checkpoint = tutils.init_checkpoint_callback(logger) - trainer_options = dict( + # fit model + trainer = Trainer( progress_bar_refresh_rate=0, max_epochs=1, train_percent_check=0.4, @@ -169,9 +168,6 @@ def test_running_test_no_val(tmpdir): logger=logger, early_stop_callback=False ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) assert result == 1, 'training failed to complete' @@ -238,16 +234,13 @@ def test_simple_cpu(tmpdir): hparams = tutils.get_default_hparams() model = LightningTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.1, ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) # traning complete @@ -340,15 +333,6 @@ def test_tbptt_cpu_model(tmpdir): sampler=None, ) - trainer_options = dict( - default_root_dir=tmpdir, - max_epochs=1, - truncated_bptt_steps=truncated_bptt_steps, - val_percent_check=0, - weights_summary=None, - early_stop_callback=False - ) - hparams = tutils.get_default_hparams() hparams.batch_size = batch_size hparams.in_features = truncated_bptt_steps @@ -358,7 +342,14 @@ def test_tbptt_cpu_model(tmpdir): model = BpttTestModel(hparams) # fit model - trainer = Trainer(**trainer_options) + trainer = Trainer( + default_root_dir=tmpdir, + max_epochs=1, + truncated_bptt_steps=truncated_bptt_steps, + val_percent_check=0, + weights_summary=None, + early_stop_callback=False + ) result = trainer.fit(model) assert result == 1, 'training failed to complete' diff --git a/tests/models/test_gpu.py b/tests/models/test_gpu.py index 69580fdd..dcd90b08 100644 --- a/tests/models/test_gpu.py +++ b/tests/models/test_gpu.py @@ -50,19 +50,19 @@ def test_ddp_all_dataloaders_passed_to_fit(tmpdir): tutils.set_random_master_port() model, hparams = tutils.get_default_model() - trainer_options = dict(default_root_dir=tmpdir, - progress_bar_refresh_rate=0, - max_epochs=1, - train_percent_check=0.4, - val_percent_check=0.2, - gpus=[0, 1], - distributed_backend='ddp') - fit_options = dict(train_dataloader=model.train_dataloader(), - val_dataloaders=model.val_dataloader()) - - trainer = Trainer(**trainer_options) - result = trainer.fit(model, **fit_options) + trainer = Trainer( + default_root_dir=tmpdir, + progress_bar_refresh_rate=0, + max_epochs=1, + train_percent_check=0.4, + val_percent_check=0.2, + gpus=[0, 1], + distributed_backend='ddp' + ) + result = trainer.fit(model, + train_dataloader=model.train_dataloader(), + val_dataloaders=model.val_dataloader()) assert result == 1, "DDP doesn't work with dataloaders passed to fit()." @@ -77,14 +77,12 @@ def test_cpu_slurm_save_load(tmpdir): logger = tutils.get_default_logger(tmpdir) version = logger.version - trainer_options = dict( + # fit model + trainer = Trainer( max_epochs=1, logger=logger, checkpoint_callback=ModelCheckpoint(tmpdir) ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) real_global_step = trainer.global_step @@ -115,12 +113,11 @@ def test_cpu_slurm_save_load(tmpdir): # new logger file to get meta logger = tutils.get_default_logger(tmpdir, version=version) - trainer_options = dict( + trainer = Trainer( max_epochs=1, logger=logger, checkpoint_callback=ModelCheckpoint(tmpdir), ) - trainer = Trainer(**trainer_options) model = LightningTestModel(hparams) # set the epoch start hook so we can predict before the model does the full training diff --git a/tests/models/test_restore.py b/tests/models/test_restore.py index 1165da6f..0921a3a8 100644 --- a/tests/models/test_restore.py +++ b/tests/models/test_restore.py @@ -163,12 +163,6 @@ def test_dp_resume(tmpdir): hparams = tutils.get_default_hparams() model = LightningTestModel(hparams) - trainer_options = dict( - max_epochs=1, - gpus=2, - distributed_backend='dp', - ) - # get logger logger = tutils.get_default_logger(tmpdir) @@ -176,9 +170,13 @@ def test_dp_resume(tmpdir): # logger file to get weights checkpoint = tutils.init_checkpoint_callback(logger) - # add these to the trainer options - trainer_options['logger'] = logger - trainer_options['checkpoint_callback'] = checkpoint + trainer_options = dict( + max_epochs=1, + gpus=2, + distributed_backend='dp', + logger=logger, + checkpoint_callback=checkpoint, + ) # fit model trainer = Trainer(**trainer_options) @@ -199,11 +197,13 @@ def test_dp_resume(tmpdir): # init new trainer new_logger = tutils.get_default_logger(tmpdir, version=logger.version) - trainer_options['logger'] = new_logger - trainer_options['checkpoint_callback'] = ModelCheckpoint(tmpdir) - trainer_options['train_percent_check'] = 0.5 - trainer_options['val_percent_check'] = 0.2 - trainer_options['max_epochs'] = 1 + trainer_options.update( + logger=new_logger, + checkpoint_callback=ModelCheckpoint(tmpdir), + train_percent_check=0.5, + val_percent_check=0.2, + max_epochs=1, + ) new_trainer = Trainer(**trainer_options) # set the epoch start hook so we can predict before the model does the full training @@ -240,14 +240,12 @@ def test_model_saving_loading(tmpdir): # logger file to get meta logger = tutils.get_default_logger(tmpdir) - trainer_options = dict( + # fit model + trainer = Trainer( max_epochs=1, logger=logger, checkpoint_callback=ModelCheckpoint(tmpdir) ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) # traning complete @@ -289,7 +287,8 @@ def test_model_saving_loading(tmpdir): def test_load_model_with_missing_hparams(tmpdir): - trainer_options = dict( + # fit model + trainer = Trainer( progress_bar_refresh_rate=0, max_epochs=1, checkpoint_callback=ModelCheckpoint(tmpdir, save_top_k=-1), @@ -297,9 +296,6 @@ def test_load_model_with_missing_hparams(tmpdir): default_root_dir=tmpdir, ) - # fit model - trainer = Trainer(**trainer_options) - model = LightningTestModelWithoutHyperparametersArg() trainer.fit(model) last_checkpoint = sorted(glob.glob(os.path.join(trainer.checkpoint_callback.dirpath, "*.ckpt")))[-1] diff --git a/tests/trainer/test_checks.py b/tests/trainer/test_checks.py index 2ee80377..d69ec1e6 100755 --- a/tests/trainer/test_checks.py +++ b/tests/trainer/test_checks.py @@ -21,8 +21,7 @@ def test_error_on_no_train_step(tmpdir): def forward(self, x): pass - trainer_options = dict(default_root_dir=tmpdir, max_epochs=1) - trainer = Trainer(**trainer_options) + trainer = Trainer(default_root_dir=tmpdir, max_epochs=1) with pytest.raises(MisconfigurationException): model = CurrentTestModel() @@ -37,8 +36,7 @@ def test_error_on_no_train_dataloader(tmpdir): class CurrentTestModel(TestModelBase): pass - trainer_options = dict(default_root_dir=tmpdir, max_epochs=1) - trainer = Trainer(**trainer_options) + trainer = Trainer(default_root_dir=tmpdir, max_epochs=1) with pytest.raises(MisconfigurationException): model = CurrentTestModel(hparams) @@ -56,8 +54,7 @@ def test_error_on_no_configure_optimizers(tmpdir): def training_step(self, batch, batch_idx, optimizer_idx=None): pass - trainer_options = dict(default_root_dir=tmpdir, max_epochs=1) - trainer = Trainer(**trainer_options) + trainer = Trainer(default_root_dir=tmpdir, max_epochs=1) with pytest.raises(MisconfigurationException): model = CurrentTestModel() @@ -74,8 +71,7 @@ def test_warning_on_wrong_validation_settings(tmpdir): tutils.reset_seed() hparams = tutils.get_default_hparams() - trainer_options = dict(default_root_dir=tmpdir, max_epochs=1) - trainer = Trainer(**trainer_options) + trainer = Trainer(default_root_dir=tmpdir, max_epochs=1) class CurrentTestModel(LightTrainDataloader, LightValidationDataloader, diff --git a/tests/trainer/test_dataloaders.py b/tests/trainer/test_dataloaders.py index d52e61ea..83ff481d 100644 --- a/tests/trainer/test_dataloaders.py +++ b/tests/trainer/test_dataloaders.py @@ -24,7 +24,13 @@ from tests.base import ( ) -def test_dataloader_config_errors(tmpdir): +@pytest.mark.parametrize("dataloader_options", [ + dict(train_percent_check=-0.1), + dict(train_percent_check=1.1), + dict(val_check_interval=1.1), + dict(val_check_interval=10000), +]) +def test_dataloader_config_errors(tmpdir, dataloader_options): tutils.reset_seed() class CurrentTestModel( @@ -36,63 +42,13 @@ def test_dataloader_config_errors(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # percent check < 0 - - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, - train_percent_check=-0.1, + **dataloader_options, ) - # fit model - trainer = Trainer(**trainer_options) - - with pytest.raises(ValueError): - trainer.fit(model) - - # percent check > 1 - - # logger file to get meta - trainer_options = dict( - default_root_dir=tmpdir, - max_epochs=1, - train_percent_check=1.1, - ) - - # fit model - trainer = Trainer(**trainer_options) - - with pytest.raises(ValueError): - trainer.fit(model) - - # int val_check_interval > num batches - - # logger file to get meta - trainer_options = dict( - default_root_dir=tmpdir, - max_epochs=1, - val_check_interval=10000 - ) - - # fit model - trainer = Trainer(**trainer_options) - - with pytest.raises(ValueError): - trainer.fit(model) - - # float val_check_interval > 1 - - # logger file to get meta - trainer_options = dict( - default_root_dir=tmpdir, - max_epochs=1, - val_check_interval=1.1 - ) - - # fit model - trainer = Trainer(**trainer_options) - with pytest.raises(ValueError): trainer.fit(model) @@ -111,16 +67,13 @@ def test_multiple_val_dataloader(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=1.0, ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) # verify training completed @@ -150,16 +103,13 @@ def test_multiple_test_dataloader(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # fit model - trainer = Trainer(**trainer_options) trainer.fit(model) trainer.test() @@ -184,19 +134,15 @@ def test_train_dataloaders_passed_to_fit(tmpdir): hparams = tutils.get_default_hparams() - # logger file to get meta - trainer_options = dict( + # only train passed to fit + model = CurrentTestModel(hparams) + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # only train passed to fit - model = CurrentTestModel(hparams) - trainer = Trainer(**trainer_options) - fit_options = dict(train_dataloader=model._dataloader(train=True)) - result = trainer.fit(model, **fit_options) + result = trainer.fit(model, train_dataloader=model._dataloader(train=True)) assert result == 1 @@ -214,21 +160,17 @@ def test_train_val_dataloaders_passed_to_fit(tmpdir): hparams = tutils.get_default_hparams() - # logger file to get meta - trainer_options = dict( + # train, val passed to fit + model = CurrentTestModel(hparams) + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # train, val passed to fit - model = CurrentTestModel(hparams) - trainer = Trainer(**trainer_options) - fit_options = dict(train_dataloader=model._dataloader(train=True), - val_dataloaders=model._dataloader(train=False)) - - result = trainer.fit(model, **fit_options) + result = trainer.fit(model, + train_dataloader=model._dataloader(train=True), + val_dataloaders=model._dataloader(train=False)) assert result == 1 assert len(trainer.val_dataloaders) == 1, \ f'`val_dataloaders` not initiated properly, got {trainer.val_dataloaders}' @@ -249,24 +191,20 @@ def test_all_dataloaders_passed_to_fit(tmpdir): hparams = tutils.get_default_hparams() - # logger file to get meta - trainer_options = dict( + # train, val and test passed to fit + model = CurrentTestModel(hparams) + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - # train, val and test passed to fit - model = CurrentTestModel(hparams) - trainer = Trainer(**trainer_options) - fit_options = dict(train_dataloader=model._dataloader(train=True), - val_dataloaders=model._dataloader(train=False)) - test_options = dict(test_dataloaders=model._dataloader(train=False)) + result = trainer.fit(model, + train_dataloader=model._dataloader(train=True), + val_dataloaders=model._dataloader(train=False)) - result = trainer.fit(model, **fit_options) - - trainer.test(**test_options) + trainer.test(test_dataloaders=model._dataloader(train=False)) assert result == 1 assert len(trainer.val_dataloaders) == 1, \ @@ -288,25 +226,23 @@ def test_multiple_dataloaders_passed_to_fit(tmpdir): hparams = tutils.get_default_hparams() - # logger file to get meta - trainer_options = dict( + # train, multiple val and multiple test passed to fit + model = CurrentTestModel(hparams) + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - # train, multiple val and multiple test passed to fit - model = CurrentTestModel(hparams) - trainer = Trainer(**trainer_options) - fit_options = dict(train_dataloader=model._dataloader(train=True), - val_dataloaders=[model._dataloader(train=False), - model._dataloader(train=False)]) - test_options = dict(test_dataloaders=[model._dataloader(train=False), - model._dataloader(train=False)]) + results = trainer.fit( + model, + train_dataloader=model._dataloader(train=True), + val_dataloaders=[model._dataloader(train=False), model._dataloader(train=False)], + ) + assert results - results = trainer.fit(model, **fit_options) - trainer.test(**test_options) + trainer.test(test_dataloaders=[model._dataloader(train=False), model._dataloader(train=False)]) assert len(trainer.val_dataloaders) == 2, \ f'Multiple `val_dataloaders` not initiated properly, got {trainer.val_dataloaders}' @@ -329,7 +265,6 @@ def test_mixing_of_dataloader_options(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta trainer_options = dict( default_root_dir=tmpdir, max_epochs=1, @@ -341,6 +276,7 @@ def test_mixing_of_dataloader_options(tmpdir): trainer = Trainer(**trainer_options) fit_options = dict(val_dataloaders=model._dataloader(train=False)) results = trainer.fit(model, **fit_options) + assert results # fit model trainer = Trainer(**trainer_options) @@ -506,20 +442,17 @@ def test_warning_with_few_workers(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + fit_options = dict(train_dataloader=model._dataloader(train=True), + val_dataloaders=model._dataloader(train=False)) + test_options = dict(test_dataloaders=model._dataloader(train=False)) + + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - fit_options = dict(train_dataloader=model._dataloader(train=True), - val_dataloaders=model._dataloader(train=False)) - test_options = dict(test_dataloaders=model._dataloader(train=False)) - - trainer = Trainer(**trainer_options) - # fit model with pytest.warns(UserWarning, match='train'): trainer.fit(model, **fit_options) diff --git a/tests/trainer/test_optimizers.py b/tests/trainer/test_optimizers.py index 6ac9da23..b445dcb2 100644 --- a/tests/trainer/test_optimizers.py +++ b/tests/trainer/test_optimizers.py @@ -29,16 +29,13 @@ def test_optimizer_with_scheduling(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # fit model - trainer = Trainer(**trainer_options) results = trainer.fit(model) init_lr = hparams.learning_rate @@ -68,16 +65,13 @@ def test_multi_optimizer_with_scheduling(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # fit model - trainer = Trainer(**trainer_options) results = trainer.fit(model) init_lr = hparams.learning_rate @@ -111,16 +105,13 @@ def test_multi_optimizer_with_scheduling_stepping(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # fit model - trainer = Trainer(**trainer_options) results = trainer.fit(model) init_lr = hparams.learning_rate @@ -160,17 +151,15 @@ def test_reduce_lr_on_plateau_scheduling(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # fit model - trainer = Trainer(**trainer_options) results = trainer.fit(model) + assert results assert trainer.lr_schedulers[0] == \ dict(scheduler=trainer.lr_schedulers[0]['scheduler'], monitor='val_loss', @@ -260,16 +249,13 @@ def test_none_optimizer(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, val_percent_check=0.1, train_percent_check=0.2 ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) # verify training completed @@ -291,9 +277,7 @@ def test_configure_optimizer_from_dict(tmpdir): hparams = tutils.get_default_hparams() model = CurrentTestModel(hparams) - trainer_options = dict(default_save_path=tmpdir, max_epochs=1) - # fit model - trainer = Trainer(**trainer_options) + trainer = Trainer(default_save_path=tmpdir, max_epochs=1) result = trainer.fit(model) assert result == 1 diff --git a/tests/trainer/test_trainer.py b/tests/trainer/test_trainer.py index cb650fd8..b7344a70 100644 --- a/tests/trainer/test_trainer.py +++ b/tests/trainer/test_trainer.py @@ -36,16 +36,12 @@ def test_model_pickle(tmpdir): def test_hparams_save_load(tmpdir): model = DictHparamsModel({'in_features': 28 * 28, 'out_features': 10, 'failed_key': lambda x: x}) - # logger file to get meta - trainer_options = dict( + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, ) - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) - assert result == 1 # try to load the model now @@ -69,16 +65,13 @@ def test_no_val_module(tmpdir): # logger file to get meta logger = tutils.get_default_logger(tmpdir) - trainer_options = dict( + trainer = Trainer( max_epochs=1, logger=logger, checkpoint_callback=ModelCheckpoint(tmpdir) ) - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) - # training complete assert result == 1, 'amp + ddp model failed to complete' @@ -110,14 +103,12 @@ def test_no_val_end_module(tmpdir): # logger file to get meta logger = tutils.get_default_logger(tmpdir) - trainer_options = dict( + # fit model + trainer = Trainer( max_epochs=1, logger=logger, checkpoint_callback=ModelCheckpoint(tmpdir) ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) # traning complete @@ -353,8 +344,8 @@ def test_resume_from_checkpoint(tmpdir): val_check_interval=1., ) - # fit model trainer = Trainer(**trainer_options) + # fit model trainer.fit(model) training_batches = trainer.num_training_batches @@ -399,11 +390,11 @@ def test_trainer_max_steps_and_epochs(tmpdir): model, trainer_options, num_train_samples = _init_steps_model() # define less train steps than epochs - trainer_options.update(dict( + trainer_options.update( default_root_dir=tmpdir, max_epochs=3, max_steps=num_train_samples + 10 - )) + ) # fit model trainer = Trainer(**trainer_options) @@ -414,10 +405,10 @@ def test_trainer_max_steps_and_epochs(tmpdir): assert trainer.global_step == trainer.max_steps, "Model did not stop at max_steps" # define less train epochs than steps - trainer_options.update(dict( + trainer_options.update( max_epochs=2, max_steps=trainer_options['max_epochs'] * 2 * num_train_samples - )) + ) # fit model trainer = Trainer(**trainer_options) @@ -434,13 +425,13 @@ def test_trainer_min_steps_and_epochs(tmpdir): model, trainer_options, num_train_samples = _init_steps_model() # define callback for stopping the model and default epochs - trainer_options.update(dict( + trainer_options.update( default_root_dir=tmpdir, early_stop_callback=EarlyStopping(monitor='val_loss', min_delta=1.0), val_check_interval=2, min_epochs=1, max_epochs=5 - )) + ) # define less min steps than 1 epoch trainer_options['min_steps'] = math.floor(num_train_samples / 2) @@ -484,15 +475,12 @@ def test_benchmark_option(tmpdir): # verify torch.backends.cudnn.benchmark is not turned on assert not torch.backends.cudnn.benchmark - # logger file to get meta - trainer_options = dict( + # fit model + trainer = Trainer( default_root_dir=tmpdir, max_epochs=1, benchmark=True, ) - - # fit model - trainer = Trainer(**trainer_options) result = trainer.fit(model) # verify training completed @@ -660,17 +648,15 @@ def test_trainer_interrupted_flag(tmpdir): interrupt_callback = InterruptCallback() - trainer_options = { - 'callbacks': [interrupt_callback], - 'max_epochs': 1, - 'val_percent_check': 0.1, - 'train_percent_check': 0.2, - 'progress_bar_refresh_rate': 0, - 'logger': False, - 'default_root_dir': tmpdir, - } - - trainer = Trainer(**trainer_options) + trainer = Trainer( + callbacks=[interrupt_callback], + max_epochs=1, + val_percent_check=0.1, + train_percent_check=0.2, + progress_bar_refresh_rate=0, + logger=False, + default_root_dir=tmpdir, + ) assert not trainer.interrupted trainer.fit(model) assert trainer.interrupted