MCPcopy Create free account
hub / github.com/rushter/MLAlgorithms / load

Function load

mla/datasets/base.py:12–47  ·  view source on GitHub ↗
(dataset="training", digits=np.arange(10))

Source from the content-addressed store, hash-verified

10
11def 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")

Callers 1

load_mnistFunction · 0.85

Calls 1

get_filenameFunction · 0.85

Tested by

no test coverage detected