(dataset="training", digits=np.arange(10))
| 10 | |
| 11 | def load_mnist(): |
| 12 | def load(dataset="training", digits=np.arange(10)): |
| 13 | import struct |
| 14 | from array import array as pyarray |
| 15 | from numpy import array, int8, uint8, zeros |
| 16 | |
| 17 | if dataset == "train": |
| 18 | fname_img = get_filename("data/mnist/train-images-idx3-ubyte") |
| 19 | fname_lbl = get_filename("data/mnist/train-labels-idx1-ubyte") |
| 20 | elif dataset == "test": |
| 21 | fname_img = get_filename("data/mnist/t10k-images-idx3-ubyte") |
| 22 | fname_lbl = get_filename("data/mnist/t10k-labels-idx1-ubyte") |
| 23 | else: |
| 24 | raise ValueError("Unexpected dataset name: %r" % dataset) |
| 25 | |
| 26 | flbl = open(fname_lbl, "rb") |
| 27 | magic_nr, size = struct.unpack(">II", flbl.read(8)) |
| 28 | lbl = pyarray("b", flbl.read()) |
| 29 | flbl.close() |
| 30 | |
| 31 | fimg = open(fname_img, "rb") |
| 32 | magic_nr, size, rows, cols = struct.unpack(">IIII", fimg.read(16)) |
| 33 | img = pyarray("B", fimg.read()) |
| 34 | fimg.close() |
| 35 | |
| 36 | ind = [k for k in range(size) if lbl[k] in digits] |
| 37 | N = len(ind) |
| 38 | |
| 39 | images = zeros((N, rows, cols), dtype=uint8) |
| 40 | labels = zeros((N, 1), dtype=int8) |
| 41 | for i in range(len(ind)): |
| 42 | images[i] = array( |
| 43 | img[ind[i] * rows * cols : (ind[i] + 1) * rows * cols] |
| 44 | ).reshape((rows, cols)) |
| 45 | labels[i] = lbl[ind[i]] |
| 46 | |
| 47 | return images, labels |
| 48 | |
| 49 | X_train, y_train = load("train") |
| 50 | X_test, y_test = load("test") |
no test coverage detected