mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
move unnecessary dict trainer_options (#1469)
* move unnecessary dict trainer_options * fix tests * fix tests * formatting * missing
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user