mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
updated docs
This commit is contained in:
@@ -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
|
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
Reference in New Issue
Block a user