From 5d5968033f7c2477b2c8c19933ca6ed2be3adbd9 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 12 Aug 2019 16:07:42 -0400 Subject: [PATCH] LR scheduler + train refactor (#103) * split __train up for clarity * split __train up for clarity * added lr scheduler after epoch completes --- pytorch_lightning/models/trainer.py | 174 ++++++++++++++-------------- 1 file changed, 88 insertions(+), 86 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 576f11e4..0ded3dd3 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -758,25 +758,16 @@ We recommend you switch to ddp if you want to use amp # --------------------------- # CORE TRAINING LOOP # --------------------------- - self.__train() def __train(self): # run all epochs for epoch_nb in range(self.current_epoch, self.max_nb_epochs): - # update the lr scheduler - if self.lr_schedulers is not None: - for lr_scheduler in self.lr_schedulers: - lr_scheduler.step() - + # get model model = self.__get_model() + + # update training progress in trainer and model model.current_epoch = epoch_nb - - # hook - if self.__is_function_implemented('on_epoch_start'): - model = self.__get_model() - model.on_epoch_start() - self.current_epoch = epoch_nb self.total_batches = self.nb_tng_batches + self.nb_val_batches self.batch_loss_value = 0 # accumulated grads @@ -786,92 +777,103 @@ We recommend you switch to ddp if you want to use amp self.prog_bar = tqdm.tqdm(range(self.total_batches), position=self.process_position) - for batch_nb, data_batch in enumerate(self.tng_dataloader): - self.batch_nb = batch_nb - self.global_step += 1 + # ----------------- + # RUN TNG EPOCH + # ----------------- + self.run_tng_epoch() - model = self.__get_model() - model.global_step = self.global_step - - # stop when the flag is changed or we've gone past the amount - # requested in the batches - self.total_batch_nb += 1 - met_batch_limit = batch_nb > self.nb_tng_batches - if met_batch_limit: - break - - # --------------- - # RUN TRAIN STEP - # --------------- - batch_result = self.__run_tng_batch(data_batch, batch_nb) - early_stop_epoch = batch_result == -1 - - # --------------- - # RUN VAL STEP - # --------------- - is_val_check_batch = (batch_nb + 1) % self.val_check_batch == 0 - if self.fast_dev_run or is_val_check_batch or early_stop_epoch: - self.__run_validation() - - # when batch should be saved - if (batch_nb + 1) % self.log_save_interval == 0 or early_stop_epoch: - if self.proc_rank == 0 and self.experiment is not None: - self.experiment.save() - - # when metrics should be logged - if batch_nb % self.add_log_row_interval == 0 or early_stop_epoch: - # count items in memory - # nb_params, nb_tensors = count_mem_items() - - model = self.__get_model() - metrics = self.__tng_tqdm_dic - - # add gpu memory - if self.on_gpu: - mem_map = get_gpu_memory_map() - metrics.update(mem_map) - - # add norms - if self.track_grad_norm > 0: - model = self.__get_model() - grad_norm_dic = model.grad_norm(self.track_grad_norm) - metrics.update(grad_norm_dic) - - if self.__is_function_implemented('on_tng_metrics'): - model.on_tng_metrics(metrics) - - # log metrics - scalar_metrics = self.__metrics_to_scalars( - metrics, blacklist=self.__log_vals_blacklist()) - if self.proc_rank == 0 and self.experiment is not None: - self.experiment.log(scalar_metrics, global_step=self.global_step) - self.experiment.save() - - # hook - if self.__is_function_implemented('on_batch_end'): - model = self.__get_model() - model.on_batch_end() - - # end epoch early - if early_stop_epoch: - break - - # hook - if self.__is_function_implemented('on_epoch_end'): - model = self.__get_model() - model.on_epoch_end() + # update LR schedulers + if self.lr_schedulers is not None: + for lr_scheduler in self.lr_schedulers: + lr_scheduler.step() # early stopping met_min_epochs = epoch_nb > self.min_nb_epochs if self.enable_early_stop and met_min_epochs: should_stop = self.early_stop_callback.on_epoch_end(epoch=epoch_nb, logs=self.__tng_tqdm_dic) - # stop training stop = should_stop and met_min_epochs if stop: return + def run_tng_epoch(self): + # before epoch hook + if self.__is_function_implemented('on_epoch_start'): + model = self.__get_model() + model.on_epoch_start() + + # run epoch + for batch_nb, data_batch in enumerate(self.tng_dataloader): + self.batch_nb = batch_nb + self.global_step += 1 + + model = self.__get_model() + model.global_step = self.global_step + + # stop when the flag is changed or we've gone past the amount + # requested in the batches + self.total_batch_nb += 1 + met_batch_limit = batch_nb > self.nb_tng_batches + if met_batch_limit: + break + + # --------------- + # RUN TRAIN STEP + # --------------- + batch_result = self.__run_tng_batch(data_batch, batch_nb) + early_stop_epoch = batch_result == -1 + + # --------------- + # RUN VAL STEP + # --------------- + is_val_check_batch = (batch_nb + 1) % self.val_check_batch == 0 + if self.fast_dev_run or is_val_check_batch or early_stop_epoch: + self.__run_validation() + + # when batch should be saved + if (batch_nb + 1) % self.log_save_interval == 0 or early_stop_epoch: + if self.proc_rank == 0 and self.experiment is not None: + self.experiment.save() + + # when metrics should be logged + if batch_nb % self.add_log_row_interval == 0 or early_stop_epoch: + # count items in memory + # nb_params, nb_tensors = count_mem_items() + + model = self.__get_model() + metrics = self.__tng_tqdm_dic + + # add gpu memory + if self.on_gpu: + mem_map = get_gpu_memory_map() + metrics.update(mem_map) + + # add norms + if self.track_grad_norm > 0: + model = self.__get_model() + grad_norm_dic = model.grad_norm(self.track_grad_norm) + metrics.update(grad_norm_dic) + + if self.__is_function_implemented('on_tng_metrics'): + model.on_tng_metrics(metrics) + + # log metrics + scalar_metrics = self.__metrics_to_scalars( + metrics, blacklist=self.__log_vals_blacklist()) + if self.proc_rank == 0 and self.experiment is not None: + self.experiment.log(scalar_metrics, global_step=self.global_step) + self.experiment.save() + + # end epoch early + if early_stop_epoch: + break + + # epoch end hook + if self.__is_function_implemented('on_epoch_end'): + model = self.__get_model() + model.on_epoch_end() + def __metrics_to_scalars(self, metrics, blacklist=set()): new_metrics = {} for k, v in metrics.items():