mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Add MNIST dataset & drop torchvision dep. from tests (#986)
* added custom mnist without torchvision dep * move files so it does not conflict with mnist gitignore * mock torchvision for tests * fix line too long * fix line too long * fix "module level import not at top of file" warning * move mock imports to __init__.py * simplify MNIST a lot and download directly the .pt files * further simplify and clean up mnist * revert import overrides * make as before * drop PIL requirement * move mnist.py to datasets subfolder * use logging instead of print * choose same name as in torchvision * remove torchvision and pillow also from yml file * refactor if train Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * capitalized class attr * moved mnist to models * re-added datsets ignore * better name for file variable * Update mnist.py * move dataset classes to datasets.py * new line * update * update * fix automerge * move to base folder * adapt testingmnist to new mnist base class * remove temporal fix * fix datatype * remove old testingmnist * readable * fix import * fix whitespace * docstring Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Update tests/base/datasets.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * changelog * added types * Update CHANGELOG.md Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * exist->isfile Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * index -> idx * temporary fix for trains error * better changelog message Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com>
This commit is contained in:
co-authored by
Jirka Borovec
parent
18d055a390
commit
b7de42f70d
+5
-33
@@ -7,8 +7,8 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import optim
|
||||
from torch.utils.data import DataLoader
|
||||
from torchvision import transforms
|
||||
from torchvision.datasets import MNIST
|
||||
|
||||
from tests.base.datasets import TestingMNIST
|
||||
|
||||
try:
|
||||
from test_tube import HyperOptArgumentParser
|
||||
@@ -18,29 +18,6 @@ except ImportError:
|
||||
|
||||
from pytorch_lightning.core.lightning import LightningModule
|
||||
|
||||
# TODO: remove after getting own MNIST
|
||||
# TEMPORAL FIX, https://github.com/pytorch/vision/issues/1938
|
||||
import urllib.request
|
||||
opener = urllib.request.build_opener()
|
||||
opener.addheaders = [('User-agent', 'Mozilla/5.0')]
|
||||
urllib.request.install_opener(opener)
|
||||
|
||||
|
||||
class TestingMNIST(MNIST):
|
||||
|
||||
def __init__(self, root, train=True, transform=None, target_transform=None,
|
||||
download=False, num_samples=8000):
|
||||
super().__init__(
|
||||
root,
|
||||
train=train,
|
||||
transform=transform,
|
||||
target_transform=target_transform,
|
||||
download=download
|
||||
)
|
||||
# take just a subset of MNIST dataset
|
||||
self.data = self.data[:num_samples]
|
||||
self.targets = self.targets[:num_samples]
|
||||
|
||||
|
||||
class DictHparamsModel(LightningModule):
|
||||
|
||||
@@ -61,8 +38,7 @@ class DictHparamsModel(LightningModule):
|
||||
return torch.optim.Adam(self.parameters(), lr=0.02)
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(TestingMNIST(os.getcwd(), train=True, download=True,
|
||||
transform=transforms.ToTensor()), batch_size=32)
|
||||
return DataLoader(TestingMNIST(os.getcwd(), train=True, download=True), batch_size=32)
|
||||
|
||||
|
||||
class TestModelBase(LightningModule):
|
||||
@@ -178,17 +154,13 @@ class TestModelBase(LightningModule):
|
||||
return [optimizer], [scheduler]
|
||||
|
||||
def prepare_data(self):
|
||||
transform = transforms.Compose([transforms.ToTensor(),
|
||||
transforms.Normalize((0.5,), (1.0,))])
|
||||
_ = TestingMNIST(root=self.hparams.data_root, train=True,
|
||||
transform=transform, download=True, num_samples=2000)
|
||||
download=True, num_samples=2000)
|
||||
|
||||
def _dataloader(self, train):
|
||||
# init data generators
|
||||
transform = transforms.Compose([transforms.ToTensor(),
|
||||
transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = TestingMNIST(root=self.hparams.data_root, train=train,
|
||||
transform=transform, download=False, num_samples=2000)
|
||||
download=False, num_samples=2000)
|
||||
|
||||
# when using multi-node we need to add the datasampler
|
||||
batch_size = self.hparams.batch_size
|
||||
|
||||
Reference in New Issue
Block a user