From 06928f721f157bca24ea32cff056884f796dd7b4 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 8 Jul 2019 10:51:31 -0400 Subject: [PATCH] testing slurm ddp --- .../trainer_gpu_cluster_template.py | 19 +++++-------------- 1 file changed, 5 insertions(+), 14 deletions(-) diff --git a/examples/new_project_templates/trainer_gpu_cluster_template.py b/examples/new_project_templates/trainer_gpu_cluster_template.py index 3e33d95d..2e029b7e 100644 --- a/examples/new_project_templates/trainer_gpu_cluster_template.py +++ b/examples/new_project_templates/trainer_gpu_cluster_template.py @@ -42,21 +42,13 @@ def main(hparams, cluster, results_dict): :param hparams: :return: """ - on_gpu = hparams.gpus is not None and torch.cuda.is_available() - - device = 'cuda' if on_gpu else 'cpu' - hparams.__setattr__('device', device) - hparams.__setattr__('on_gpu', on_gpu) - hparams.__setattr__('nb_gpus', torch.cuda.device_count()) - hparams.__setattr__('inference_mode', hparams.model_load_weights_path is not None) - # delay each training start to not overwrite logs - process_position, current_gpu = TRAINING_MODEL.get_process_position(hparams.gpus) + process_position, current_gpu = LightningTemplateModel.get_process_position(hparams.gpus) sleep(process_position + 1) # init experiment log_dir = os.path.dirname(os.path.realpath(__file__)) - log_dir = os.path.join(log_dir, 'test_tube_demo_logs') + log_dir = os.path.join(log_dir, 'pt_lightning_demo_logs') exp = Experiment( name='test_tube_exp', save_dir=log_dir, @@ -69,7 +61,7 @@ def main(hparams, cluster, results_dict): # build model print('loading model...') - model = TRAINING_MODEL(hparams) + model = LightningTemplateModel(hparams) print('model built') # callbacks @@ -91,7 +83,7 @@ def main(hparams, cluster, results_dict): # configure trainer trainer = Trainer( - experiment=exp, + experiment=hyperparams.tt_name, cluster=cluster, checkpoint_callback=checkpoint, early_stop_callback=early_stop, @@ -105,9 +97,8 @@ def main(hparams, cluster, results_dict): def get_default_parser(strategy, root_dir): - possible_model_names = list(AVAILABLE_MODELS.keys()) parser = HyperOptArgumentParser(strategy=strategy, add_help=False) - add_default_args(parser, root_dir, possible_model_names=possible_model_names, rand_seed=SEED) + add_default_args(parser, root_dir, rand_seed=SEED) return parser