From 34b824a9d3d0fdd377da675e0398c66ab5e16e7b Mon Sep 17 00:00:00 2001 From: Anton Konstantinov Date: Thu, 5 Sep 2019 14:13:06 +0300 Subject: [PATCH] Implement correct transfer to GPU for batches (#200) --- pytorch_lightning/models/trainer.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index abb4b167..c2fe2bcd 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -995,10 +995,13 @@ class Trainer(TrainerIO): return blacklist def transfer_batch_to_gpu(self, batch, gpu_id): - # base case - if isinstance(batch, torch.Tensor): + # base case: object can be directly moved using `cuda` or `to` + if callable(getattr(batch, 'cuda', None)): return batch.cuda(gpu_id) + elif callable(getattr(batch, 'to', None)): + return batch.to(torch.device('cuda', gpu_id)) + # when list elif isinstance(batch, list): for i, x in enumerate(batch):