* new image

* new image
This commit is contained in:
William Falcon
2020-03-06 06:33:28 -05:00
committed by GitHub
parent f996c2892f
commit d80215ed9d
2 changed files with 18 additions and 13 deletions
+3 -3
View File
@@ -494,13 +494,13 @@ In this method we do all the preparation we need to do once (instead of on every
self.test_dataset = mnist_test
def train_dataloader(self):
return DataLoader(train_dataset, batch_size=64)
return DataLoader(self.train_dataset, batch_size=64)
def val_dataloader(self):
return DataLoader(mnist_val, batch_size=64)
return DataLoader(self.mnist_val, batch_size=64)
def test_dataloader(self):
return DataLoader(mnist_test, batch_size=64)
return DataLoader(self.mnist_test, batch_size=64)
The `prepare_data` method is also a good place to do any data processing that needs to be done only
once (ie: download or tokenize, etc...).
+15 -10
View File
@@ -236,11 +236,11 @@ Data preparation
----------------
Data preparation in PyTorch follows 5 steps:
1. Download
2. Clean and (maybe) save to disk
3. Load inside dataset
4. Apply transforms (rotate, tokenize, etc...)
5. Wrap inside a dataloader
1. Download
2. Clean and (maybe) save to disk
3. Load inside dataset
4. Apply transforms (rotate, tokenize, etc...)
5. Wrap inside a dataloader
When working in distributed settings, steps 1 and 2 have to be done
from a single GPU, otherwise you will overwrite these files from
@@ -251,8 +251,10 @@ allow for this
def prepare_data(self):
# download
mnist_train = MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor())
mnist_test = MNIST(os.getcwd(), train=False, download=True, transform=transforms.ToTensor())
mnist_train = MNIST(os.getcwd(), train=True, download=True,
transform=transforms.ToTensor())
mnist_test = MNIST(os.getcwd(), train=False, download=True,
transform=transforms.ToTensor())
# train/val split
mnist_train, mnist_val = random_split(mnist_train, [55000, 5000])
@@ -263,13 +265,13 @@ allow for this
self.test_dataset = mnist_test
def train_dataloader(self):
return DataLoader(train_dataset, batch_size=64)
return DataLoader(self.train_dataset, batch_size=64)
def val_dataloader(self):
return DataLoader(mnist_val, batch_size=64)
return DataLoader(self.mnist_val, batch_size=64)
def test_dataloader(self):
return DataLoader(mnist_test, batch_size=64)
return DataLoader(self.mnist_test, batch_size=64)
.. note:: ``prepare_data`` is called once.
@@ -305,6 +307,9 @@ Check out this
`COLAB <https://colab.research.google.com/drive/1F_RNcHzTfFuQf-LeKvSlud6x7jXYkG31#scrollTo=HOk9c4_35FKg>`_
for a live demo.
LightningModule Class
---------------------
"""
from .decorators import data_loader