From 73a7cf3c9988c58c43238d4b924535cb3a53ca8e Mon Sep 17 00:00:00 2001 From: William Falcon Date: Fri, 4 Oct 2019 15:53:44 -0400 Subject: [PATCH] Mem crash (#299) * fixes memory crash * fixes memory crash --- pytorch_lightning/root_module/memory.py | 2 +- tests/test_models.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/pytorch_lightning/root_module/memory.py b/pytorch_lightning/root_module/memory.py index 73cea905..818e9643 100644 --- a/pytorch_lightning/root_module/memory.py +++ b/pytorch_lightning/root_module/memory.py @@ -35,7 +35,7 @@ class ModelSummary(object): out_sizes = [] input_ = self.model.example_input_array - if self.model.use_ddp: + if self.model.use_ddp or self.model.use_dp: input_ = input_.cuda(0) if self.model.trainer.use_amp: diff --git a/tests/test_models.py b/tests/test_models.py index cdb59f74..71e72488 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -128,8 +128,8 @@ def test_dp_resume(): dp_model = new_trainer.model dp_model.eval() - for dataloader in trainer.get_train_dataloader(): - run_prediction(dataloader, dp_model, dp=True) + dataloader = trainer.get_train_dataloader() + run_prediction(dataloader, dp_model, dp=True) # new model model = LightningTestModel(hparams)