diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 36dfcf76..f30f3173 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -366,6 +366,7 @@ class Trainer(TrainerIO): # filter out the weights that were done on gpu so we can load on good old cpus self.optimizers = model.configure_optimizers() + model.cuda() model = LightningDataParallel(model, device_ids=self.data_parallel_device_ids) # run through amp wrapper