move unnecessary dict trainer_options (#1469)

* move unnecessary dict trainer_options

* fix tests

* fix tests

* formatting

* missing
This commit is contained in:
Jirka Borovec
2020-05-01 10:43:58 -04:00
committed by GitHub
parent 97c7b6b314
commit 34bc149359
14 changed files with 161 additions and 301 deletions
-2
View File
@@ -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)
+10 -11
View File
@@ -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'
+4 -11
View File
@@ -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)
+1 -3
View File
@@ -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
+2 -5
View File
@@ -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})
+3 -5
View File
@@ -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]
+9 -15
View File
@@ -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)
+14 -23
View File
@@ -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'
+15 -18
View File
@@ -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
+18 -22
View File
@@ -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]
+4 -8
View File
@@ -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,
+47 -114
View File
@@ -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)
+12 -28
View File
@@ -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
+22 -36
View File
@@ -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