From e04dfb37fd3158a933f5da3e0a670650fc104630 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Fri, 18 Oct 2019 15:54:11 +0200 Subject: [PATCH] changes to seed for tests --- tests/test_logging.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/tests/test_logging.py b/tests/test_logging.py index 2458639b..f9a5eac6 100644 --- a/tests/test_logging.py +++ b/tests/test_logging.py @@ -10,6 +10,11 @@ from pytorch_lightning.testing import LightningTestModel from pytorch_lightning.logging import LightningLoggerBase, rank_zero_only from .test_models import get_hparams, get_test_tube_logger, init_save_dir, clear_save_dir + +RANDOM_FILE_PATHS = list(np.random.randint(12000, 19000, 1000)) +ROOT_SEED = 1234 +torch.manual_seed(ROOT_SEED) +np.random.seed(ROOT_SEED) RANDOM_SEEDS = list(np.random.randint(0, 10000, 1000)) @@ -101,7 +106,7 @@ def test_mlflow_logger(): assert result == 1, "Training failed" - n = np.random.randint(0, 10000000, 1)[0] + n = RANDOM_FILE_PATHS.pop() shutil.move(mlflow_dir, mlflow_dir + f'_{n}') @@ -136,7 +141,7 @@ def test_mlflow_pickle(): trainer2 = pickle.loads(pkl_bytes) trainer2.logger.log_metrics({"acc": 1.0}) - n = np.random.randint(0, 10000000, 1)[0] + n = RANDOM_FILE_PATHS.pop() shutil.move(mlflow_dir, mlflow_dir + f'_{n}')