updated docs

This commit is contained in:
William Falcon
2019-07-25 11:35:11 -04:00
parent 9fa8120805
commit 0f79e9d74e
@@ -248,22 +248,16 @@ Pytorch DataLoader
**Example** **Example**
``` {.python} ``` {.python}
@property @ptl.data_loader
def tng_dataloader(self): def tng_dataloader(self):
if self._tng_dataloader is None: transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
try: dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True)
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) loader = torch.utils.data.DataLoader(
dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True) dataset=dataset,
loader = torch.utils.data.DataLoader( batch_size=self.hparams.batch_size,
dataset=dataset, shuffle=True
batch_size=self.hparams.batch_size, )
shuffle=True return loader
)
self._tng_dataloader = loader
except Exception as e:
raise e
return self._tng_dataloader
``` ```
--- ---
@@ -281,22 +275,17 @@ Pytorch DataLoader
**Example** **Example**
``` {.python} ``` {.python}
@property @ptl.data_loader
def val_dataloader(self): def val_dataloader(self):
if self._val_dataloader is None: transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
try: dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) loader = torch.utils.data.DataLoader(
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True) dataset=dataset,
loader = torch.utils.data.DataLoader( batch_size=self.hparams.batch_size,
dataset=dataset, shuffle=True
batch_size=self.hparams.batch_size, )
shuffle=True
) return loader
self._val_dataloader = loader
except Exception as e:
raise e
return self._val_dataloader
``` ```
--- ---
@@ -314,22 +303,17 @@ Pytorch DataLoader
**Example** **Example**
``` {.python} ``` {.python}
@property @ptl.data_loader
def test_dataloader(self): def test_dataloader(self):
if self._test_dataloader is None: transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
try: dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) loader = torch.utils.data.DataLoader(
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True) dataset=dataset,
loader = torch.utils.data.DataLoader( batch_size=self.hparams.batch_size,
dataset=dataset, shuffle=True
batch_size=self.hparams.batch_size, )
shuffle=True
) return loader
self._test_dataloader = loader
except Exception as e:
raise e
return self._test_dataloader
``` ```
--- ---