From 62774ffacb05d8034f4c55b2e66d106ef11c2951 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Wed, 3 Jul 2019 16:17:56 -0400 Subject: [PATCH] added single node distdataparallel --- pytorch_lightning/models/trainer.py | 14 ++++++-------- 1 file changed, 6 insertions(+), 8 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 88f66db3..c3f29ced 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -250,10 +250,6 @@ class Trainer(TrainerIO): # ----------------------------- def fit(self, model): - # give model convenience properties - model.trainer = self - model.experiment = self.experiment - # transfer data loaders from model self.__get_dataloaders(model) @@ -284,7 +280,7 @@ class Trainer(TrainerIO): # when GPU is called, spawn off a single worker for each gpu if self.on_gpu: rank = 0 - # self.model = model + self.experiment = self.experiment.get_meta_copy() mp.spawn(dummy, nprocs=len(self.data_parallel_device_ids), args=(self, )) else: self.__run_pretrain_routine(model) @@ -295,7 +291,7 @@ class Trainer(TrainerIO): # del state['experiment'] return state - def __dp_train(self, gpu_nb, proc_rank): + def __dp_train(self, gpu_nb, proc_rank, model): """ Entry point into a DP thread :param gpu_nb: @@ -303,6 +299,7 @@ class Trainer(TrainerIO): :param cluster_obj: :return: """ + # TODO: pass in ip ip = "127.0.0.1" print(self.data_parallel_device_ids) @@ -321,14 +318,15 @@ class Trainer(TrainerIO): # continue training routine self.__run_pretrain_routine(model) - - def __run_pretrain_routine(self, model): """ Sanity check a few things before starting actual training :param model: :return: """ + # give model convenience properties + model.trainer = self + model.experiment = self.experiment # run tiny validation to make sure program won't crash during val _ = self.validate(model, self.val_dataloader, max_batches=self.nb_sanity_val_steps)