diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index a1a82c0c..f18ebb11 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -270,6 +270,8 @@ class Trainer(TrainerIO): # MODEL TRAINING # ----------------------------- def fit(self, model): + # set local properties on the model + self.model.on_gpu = self.on_gpu # transfer data loaders from model self.__get_dataloaders(model) diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index 969c3e64..cd36b19f 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -20,19 +20,12 @@ class LightningModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks): self.current_epoch = 0 self.global_step = 0 self.loaded_optimizer_states_dict = {} - self.fast_dev_run = hparams.fast_dev_run - self.overfit = hparams.overfit - self.gradient_clip = hparams.gradient_clip self.trainer = None self.from_lightning = True self.experiment = None # track if gpu was requested for checkpointing self.on_gpu = False - try: - self.on_gpu = hparams.on_gpu - except Exception as e: - pass # computed vars for the dataloaders self._tng_dataloader = None