mnist.py 12.9 KB
Newer Older
Tian Qi Chen's avatar
Tian Qi Chen committed
1
from __future__ import print_function
2
from .vision import VisionDataset
3
import warnings
Tian Qi Chen's avatar
Tian Qi Chen committed
4
5
6
from PIL import Image
import os
import os.path
7
import gzip
8
import numpy as np
Tian Qi Chen's avatar
Tian Qi Chen committed
9
10
import torch
import codecs
11
from .utils import download_url, makedir_exist_ok
Tian Qi Chen's avatar
Tian Qi Chen committed
12

13

14
class MNIST(VisionDataset):
15
16
17
    """`MNIST <http://yann.lecun.com/exdb/mnist/>`_ Dataset.

    Args:
18
19
        root (string): Root directory of dataset where ``MNIST/processed/training.pt``
            and  ``MNIST/processed/test.pt`` exist.
20
21
22
23
24
25
26
27
28
29
        train (bool, optional): If True, creates dataset from ``training.pt``,
            otherwise from ``test.pt``.
        download (bool, optional): If true, downloads the dataset from the internet and
            puts it in root directory. If dataset is already downloaded, it is not
            downloaded again.
        transform (callable, optional): A function/transform that  takes in an PIL image
            and returns a transformed version. E.g, ``transforms.RandomCrop``
        target_transform (callable, optional): A function/transform that takes in the
            target and transforms it.
    """
Tian Qi Chen's avatar
Tian Qi Chen committed
30
31
32
33
34
35
    urls = [
        'http://yann.lecun.com/exdb/mnist/train-images-idx3-ubyte.gz',
        'http://yann.lecun.com/exdb/mnist/train-labels-idx1-ubyte.gz',
        'http://yann.lecun.com/exdb/mnist/t10k-images-idx3-ubyte.gz',
        'http://yann.lecun.com/exdb/mnist/t10k-labels-idx1-ubyte.gz',
    ]
36
37
    training_file = 'training.pt'
    test_file = 'test.pt'
38
39
40
    classes = ['0 - zero', '1 - one', '2 - two', '3 - three', '4 - four',
               '5 - five', '6 - six', '7 - seven', '8 - eight', '9 - nine']

41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
    @property
    def train_labels(self):
        warnings.warn("train_labels has been renamed targets")
        return self.targets

    @property
    def test_labels(self):
        warnings.warn("test_labels has been renamed targets")
        return self.targets

    @property
    def train_data(self):
        warnings.warn("train_data has been renamed data")
        return self.data

    @property
    def test_data(self):
        warnings.warn("test_data has been renamed data")
        return self.data

Tian Qi Chen's avatar
Tian Qi Chen committed
61
    def __init__(self, root, train=True, transform=None, target_transform=None, download=False):
62
        super(MNIST, self).__init__(root)
Tian Qi Chen's avatar
Tian Qi Chen committed
63
64
        self.transform = transform
        self.target_transform = target_transform
65
        self.train = train  # training set or test set
Tian Qi Chen's avatar
Tian Qi Chen committed
66
67
68
69
70

        if download:
            self.download()

        if not self._check_exists():
71
72
            raise RuntimeError('Dataset not found.' +
                               ' You can use download=True to download it')
Tian Qi Chen's avatar
Tian Qi Chen committed
73
74

        if self.train:
75
            data_file = self.training_file
Tian Qi Chen's avatar
Tian Qi Chen committed
76
        else:
77
78
            data_file = self.test_file
        self.data, self.targets = torch.load(os.path.join(self.processed_folder, data_file))
Tian Qi Chen's avatar
Tian Qi Chen committed
79
80

    def __getitem__(self, index):
81
82
83
84
85
86
87
        """
        Args:
            index (int): Index

        Returns:
            tuple: (image, target) where target is index of the target class.
        """
88
        img, target = self.data[index], int(self.targets[index])
Tian Qi Chen's avatar
Tian Qi Chen committed
89
90
91
92
93
94
95
96
97
98
99
100
101
102

        # doing this so that it is consistent with all other datasets
        # to return a PIL Image
        img = Image.fromarray(img.numpy(), mode='L')

        if self.transform is not None:
            img = self.transform(img)

        if self.target_transform is not None:
            target = self.target_transform(target)

        return img, target

    def __len__(self):
103
        return len(self.data)
Tian Qi Chen's avatar
Tian Qi Chen committed
104

105
106
107
108
109
110
111
112
113
114
115
116
    @property
    def raw_folder(self):
        return os.path.join(self.root, self.__class__.__name__, 'raw')

    @property
    def processed_folder(self):
        return os.path.join(self.root, self.__class__.__name__, 'processed')

    @property
    def class_to_idx(self):
        return {_class: i for i, _class in enumerate(self.classes)}

Tian Qi Chen's avatar
Tian Qi Chen committed
117
    def _check_exists(self):
118
119
120
121
        return (os.path.exists(os.path.join(self.processed_folder,
                                            self.training_file)) and
                os.path.exists(os.path.join(self.processed_folder,
                                            self.test_file)))
122
123
124
125
126
127
128
129
130

    @staticmethod
    def extract_gzip(gzip_path, remove_finished=False):
        print('Extracting {}'.format(gzip_path))
        with open(gzip_path.replace('.gz', ''), 'wb') as out_f, \
                gzip.GzipFile(gzip_path) as zip_f:
            out_f.write(zip_f.read())
        if remove_finished:
            os.unlink(gzip_path)
Tian Qi Chen's avatar
Tian Qi Chen committed
131
132

    def download(self):
133
        """Download the MNIST data if it doesn't exist in processed_folder already."""
Tian Qi Chen's avatar
Tian Qi Chen committed
134
135
136
137

        if self._check_exists():
            return

138
139
        makedir_exist_ok(self.raw_folder)
        makedir_exist_ok(self.processed_folder)
Tian Qi Chen's avatar
Tian Qi Chen committed
140

141
        # download files
Tian Qi Chen's avatar
Tian Qi Chen committed
142
143
        for url in self.urls:
            filename = url.rpartition('/')[2]
144
            file_path = os.path.join(self.raw_folder, filename)
145
            download_url(url, root=self.raw_folder, filename=filename, md5=None)
146
            self.extract_gzip(gzip_path=file_path, remove_finished=True)
Tian Qi Chen's avatar
Tian Qi Chen committed
147
148

        # process and save as torch files
Adam Paszke's avatar
Adam Paszke committed
149
150
        print('Processing...')

Tian Qi Chen's avatar
Tian Qi Chen committed
151
        training_set = (
152
153
            read_image_file(os.path.join(self.raw_folder, 'train-images-idx3-ubyte')),
            read_label_file(os.path.join(self.raw_folder, 'train-labels-idx1-ubyte'))
Tian Qi Chen's avatar
Tian Qi Chen committed
154
155
        )
        test_set = (
156
157
            read_image_file(os.path.join(self.raw_folder, 't10k-images-idx3-ubyte')),
            read_label_file(os.path.join(self.raw_folder, 't10k-labels-idx1-ubyte'))
Tian Qi Chen's avatar
Tian Qi Chen committed
158
        )
159
        with open(os.path.join(self.processed_folder, self.training_file), 'wb') as f:
Tian Qi Chen's avatar
Tian Qi Chen committed
160
            torch.save(training_set, f)
161
        with open(os.path.join(self.processed_folder, self.test_file), 'wb') as f:
Tian Qi Chen's avatar
Tian Qi Chen committed
162
163
164
165
            torch.save(test_set, f)

        print('Done!')

166
167
    def extra_repr(self):
        return "Split: {}".format("Train" if self.train is True else "Test")
168

169

170
class FashionMNIST(MNIST):
171
172
173
    """`Fashion-MNIST <https://github.com/zalandoresearch/fashion-mnist>`_ Dataset.

    Args:
174
175
        root (string): Root directory of dataset where ``Fashion-MNIST/processed/training.pt``
            and  ``Fashion-MNIST/processed/test.pt`` exist.
176
177
178
179
180
181
182
183
184
        train (bool, optional): If True, creates dataset from ``training.pt``,
            otherwise from ``test.pt``.
        download (bool, optional): If true, downloads the dataset from the internet and
            puts it in root directory. If dataset is already downloaded, it is not
            downloaded again.
        transform (callable, optional): A function/transform that  takes in an PIL image
            and returns a transformed version. E.g, ``transforms.RandomCrop``
        target_transform (callable, optional): A function/transform that takes in the
            target and transforms it.
185
186
187
188
189
190
191
    """
    urls = [
        'http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/train-images-idx3-ubyte.gz',
        'http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/train-labels-idx1-ubyte.gz',
        'http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/t10k-images-idx3-ubyte.gz',
        'http://fashion-mnist.s3-website.eu-central-1.amazonaws.com/t10k-labels-idx1-ubyte.gz',
    ]
192
193
    classes = ['T-shirt/top', 'Trouser', 'Pullover', 'Dress', 'Coat', 'Sandal',
               'Shirt', 'Sneaker', 'Bag', 'Ankle boot']
194
195


hysts's avatar
hysts committed
196
197
198
199
class KMNIST(MNIST):
    """`Kuzushiji-MNIST <https://github.com/rois-codh/kmnist>`_ Dataset.

    Args:
200
201
        root (string): Root directory of dataset where ``KMNIST/processed/training.pt``
            and  ``KMNIST/processed/test.pt`` exist.
hysts's avatar
hysts committed
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
        train (bool, optional): If True, creates dataset from ``training.pt``,
            otherwise from ``test.pt``.
        download (bool, optional): If true, downloads the dataset from the internet and
            puts it in root directory. If dataset is already downloaded, it is not
            downloaded again.
        transform (callable, optional): A function/transform that  takes in an PIL image
            and returns a transformed version. E.g, ``transforms.RandomCrop``
        target_transform (callable, optional): A function/transform that takes in the
            target and transforms it.
    """
    urls = [
        'http://codh.rois.ac.jp/kmnist/dataset/kmnist/train-images-idx3-ubyte.gz',
        'http://codh.rois.ac.jp/kmnist/dataset/kmnist/train-labels-idx1-ubyte.gz',
        'http://codh.rois.ac.jp/kmnist/dataset/kmnist/t10k-images-idx3-ubyte.gz',
        'http://codh.rois.ac.jp/kmnist/dataset/kmnist/t10k-labels-idx1-ubyte.gz',
    ]
    classes = ['o', 'ki', 'su', 'tsu', 'na', 'ha', 'ma', 'ya', 're', 'wo']


221
class EMNIST(MNIST):
Alex Alemi's avatar
Alex Alemi committed
222
    """`EMNIST <https://www.westernsydney.edu.au/bens/home/reproducible_research/emnist>`_ Dataset.
223
224

    Args:
225
226
        root (string): Root directory of dataset where ``EMNIST/processed/training.pt``
            and  ``EMNIST/processed/test.pt`` exist.
227
228
229
230
231
232
233
234
235
236
237
238
239
        split (string): The dataset has 6 different splits: ``byclass``, ``bymerge``,
            ``balanced``, ``letters``, ``digits`` and ``mnist``. This argument specifies
            which one to use.
        train (bool, optional): If True, creates dataset from ``training.pt``,
            otherwise from ``test.pt``.
        download (bool, optional): If true, downloads the dataset from the internet and
            puts it in root directory. If dataset is already downloaded, it is not
            downloaded again.
        transform (callable, optional): A function/transform that  takes in an PIL image
            and returns a transformed version. E.g, ``transforms.RandomCrop``
        target_transform (callable, optional): A function/transform that takes in the
            target and transforms it.
    """
Alex Alemi's avatar
Alex Alemi committed
240
241
    # Updated URL from https://www.westernsydney.edu.au/bens/home/reproducible_research/emnist
    url = 'https://cloudstor.aarnet.edu.au/plus/index.php/s/54h3OuGJhFLwAlQ/download'
242
243
244
245
246
247
248
249
250
251
252
    splits = ('byclass', 'bymerge', 'balanced', 'letters', 'digits', 'mnist')

    def __init__(self, root, split, **kwargs):
        if split not in self.splits:
            raise ValueError('Split "{}" not found. Valid splits are: {}'.format(
                split, ', '.join(self.splits),
            ))
        self.split = split
        self.training_file = self._training_file(split)
        self.test_file = self._test_file(split)
        super(EMNIST, self).__init__(root, **kwargs)
Tian Qi Chen's avatar
Tian Qi Chen committed
253

254
255
    @staticmethod
    def _training_file(split):
256
257
        return 'training_{}.pt'.format(split)

258
259
    @staticmethod
    def _test_file(split):
260
261
262
263
264
265
        return 'test_{}.pt'.format(split)

    def download(self):
        """Download the EMNIST data if it doesn't exist in processed_folder already."""
        import shutil
        import zipfile
266

267
268
269
        if self._check_exists():
            return

270
271
        makedir_exist_ok(self.raw_folder)
        makedir_exist_ok(self.processed_folder)
272

273
        # download files
274
        filename = self.url.rpartition('/')[2]
275
276
        file_path = os.path.join(self.raw_folder, filename)
        download_url(self.url, root=self.raw_folder, filename=filename, md5=None)
277
278
279

        print('Extracting zip archive')
        with zipfile.ZipFile(file_path) as zip_f:
280
            zip_f.extractall(self.raw_folder)
281
        os.unlink(file_path)
282
        gzip_folder = os.path.join(self.raw_folder, 'gzip')
283
284
        for gzip_file in os.listdir(gzip_folder):
            if gzip_file.endswith('.gz'):
285
                self.extract_gzip(gzip_path=os.path.join(gzip_folder, gzip_file))
286
287
288
289
290

        # process and save as torch files
        for split in self.splits:
            print('Processing ' + split)
            training_set = (
291
292
                read_image_file(os.path.join(gzip_folder, 'emnist-{}-train-images-idx3-ubyte'.format(split))),
                read_label_file(os.path.join(gzip_folder, 'emnist-{}-train-labels-idx1-ubyte'.format(split)))
293
294
            )
            test_set = (
295
296
                read_image_file(os.path.join(gzip_folder, 'emnist-{}-test-images-idx3-ubyte'.format(split))),
                read_label_file(os.path.join(gzip_folder, 'emnist-{}-test-labels-idx1-ubyte'.format(split)))
297
            )
298
            with open(os.path.join(self.processed_folder, self._training_file(split)), 'wb') as f:
299
                torch.save(training_set, f)
300
            with open(os.path.join(self.processed_folder, self._test_file(split)), 'wb') as f:
301
                torch.save(test_set, f)
302
        shutil.rmtree(gzip_folder)
303
304
305
306
307
308

        print('Done!')


def get_int(b):
    return int(codecs.encode(b, 'hex'), 16)
Tian Qi Chen's avatar
Tian Qi Chen committed
309

310

Tian Qi Chen's avatar
Tian Qi Chen committed
311
312
313
314
315
def read_label_file(path):
    with open(path, 'rb') as f:
        data = f.read()
        assert get_int(data[:4]) == 2049
        length = get_int(data[4:8])
316
317
        parsed = np.frombuffer(data, dtype=np.uint8, offset=8)
        return torch.from_numpy(parsed).view(length).long()
Tian Qi Chen's avatar
Tian Qi Chen committed
318

319

Tian Qi Chen's avatar
Tian Qi Chen committed
320
321
322
323
324
325
326
def read_image_file(path):
    with open(path, 'rb') as f:
        data = f.read()
        assert get_int(data[:4]) == 2051
        length = get_int(data[4:8])
        num_rows = get_int(data[8:12])
        num_cols = get_int(data[12:16])
327
328
        parsed = np.frombuffer(data, dtype=np.uint8, offset=16)
        return torch.from_numpy(parsed).view(length, num_rows, num_cols)