mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Fixed tests (#340)
* removed hparam calls * removed hparam calls * removed hparam calls * removed hparam calls * removed hparam calls * Update test_models.py
This commit is contained in:
@@ -54,8 +54,10 @@ class TestTubeLogger(LightningLoggerBase):
|
||||
self.save()
|
||||
self.close()
|
||||
|
||||
@rank_zero_only
|
||||
def close(self):
|
||||
self.experiment.close()
|
||||
exp = self.experiment
|
||||
exp.close()
|
||||
|
||||
@property
|
||||
def rank(self):
|
||||
|
||||
@@ -142,7 +142,6 @@ def test_dp_resume():
|
||||
|
||||
# get logger
|
||||
logger = get_test_tube_logger(debug=False)
|
||||
logger.log_hyperparams(hparams)
|
||||
|
||||
# exp file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
@@ -218,8 +217,6 @@ def test_running_test_pretrained_model_ddp():
|
||||
|
||||
# exp file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
# exp file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
@@ -265,8 +262,6 @@ def test_running_test_after_fitting():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
@@ -308,8 +303,6 @@ def test_running_test_without_val():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
@@ -349,8 +342,6 @@ def test_running_test_pretrained_model():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
@@ -396,8 +387,6 @@ def test_running_test_pretrained_model_dp():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
@@ -648,8 +637,6 @@ def test_no_val_module():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
@@ -695,8 +682,6 @@ def test_no_val_end_module():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
@@ -831,8 +816,6 @@ def test_cpu_restore_training():
|
||||
# logger file to get meta
|
||||
test_logger_version = 10
|
||||
logger = get_test_tube_logger(False, version=test_logger_version)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=2,
|
||||
@@ -924,8 +907,6 @@ def test_cpu_slurm_save_load():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
version = logger.version
|
||||
|
||||
@@ -962,8 +943,6 @@ def test_cpu_slurm_save_load():
|
||||
|
||||
# new logger file to get meta
|
||||
logger = get_test_tube_logger(False, version=version)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
@@ -1050,8 +1029,6 @@ def test_model_saving_loading():
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
@@ -1136,8 +1113,6 @@ def test_amp_gpu_ddp_slurm_managed():
|
||||
|
||||
# exp file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
# exp file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
@@ -1358,7 +1333,6 @@ def test_ddp_sampler_error():
|
||||
model = LightningTestModel(hparams, force_remove_distributed_sampler=True)
|
||||
|
||||
logger = get_test_tube_logger(True)
|
||||
logger.save()
|
||||
|
||||
trainer = Trainer(
|
||||
logger=logger,
|
||||
@@ -1483,8 +1457,6 @@ def run_gpu_model_test(trainer_options, model, hparams, on_gpu=True):
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
|
||||
Reference in New Issue
Block a user