-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathload.go
More file actions
108 lines (96 loc) · 3.41 KB
/
Copy pathload.go
File metadata and controls
108 lines (96 loc) · 3.41 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
package main
import (
"encoding/binary"
"fmt"
"os"
)
func loadTrainingAndTestSet() (X_train, X_test [][]float32, Y_train, Y_test [][]uint8, loadErr error) {
dir := "./datasets"
var trainSetCount, testSetCount int32 = 60_000, 10_000
X_train, loadErr = loadImageFile(fmt.Sprintf("%v/train-images-idx3-ubyte", dir), trainSetCount)
if loadErr != nil {
return nil, nil, nil, nil, fmt.Errorf("unable to load training set images, %w", loadErr)
}
X_test, loadErr = loadImageFile(fmt.Sprintf("%v/t10k-images-idx3-ubyte", dir), testSetCount)
if loadErr != nil {
return nil, nil, nil, nil, fmt.Errorf("unable to load test set images, %w", loadErr)
}
Y_train, loadErr = loadLabelFile(fmt.Sprintf("%v/train-labels-idx1-ubyte", dir), trainSetCount)
if loadErr != nil {
return nil, nil, nil, nil, fmt.Errorf("unable to load training set labels, %w", loadErr)
}
Y_test, loadErr = loadLabelFile(fmt.Sprintf("%v/t10k-labels-idx1-ubyte", dir), testSetCount)
if loadErr != nil {
return nil, nil, nil, nil, fmt.Errorf("unable to load test set labels, %w", loadErr)
}
return X_train, X_test, Y_train, Y_test, nil
}
func loadLabelFile(filename string, expectedLabelCount int32) ([][]uint8, error) {
f, err := os.Open(filename)
if err != nil {
return nil, fmt.Errorf("unable to open file: %v, %w", filename, err)
}
defer f.Close()
var magicNumber int32
_ = binary.Read(f, binary.BigEndian, &magicNumber)
if magicNumber != 0x00000801 {
return nil, fmt.Errorf("invalid magic number in the file: %v", filename)
}
var labelCount int32
_ = binary.Read(f, binary.BigEndian, &labelCount)
if labelCount != expectedLabelCount {
return nil, fmt.Errorf("invalid number of labels in the file: %v", filename)
}
clasificationClassCount := 10
labels := make([][]uint8, 0, labelCount)
for range labelCount {
var label uint8
err = binary.Read(f, binary.BigEndian, &label)
if err != nil || label > 9 {
return nil, fmt.Errorf("invalid label in the file: %v", filename)
}
oneHotEncoding := make([]uint8, clasificationClassCount)
oneHotEncoding[label] = 1
labels = append(labels, oneHotEncoding)
}
return labels, nil
}
func loadImageFile(filename string, expectedImageCount int32) ([][]float32, error) {
f, err := os.Open(filename)
if err != nil {
return nil, fmt.Errorf("unable to open file: %v, %w", filename, err)
}
defer f.Close()
var magicNumber int32
_ = binary.Read(f, binary.BigEndian, &magicNumber)
if magicNumber != 0x00000803 {
return nil, fmt.Errorf("invalid magic number in the file: %v", filename)
}
var imageCount int32
_ = binary.Read(f, binary.BigEndian, &imageCount)
if imageCount != expectedImageCount {
return nil, fmt.Errorf("invalid number of images in the file: %v", filename)
}
var rowCount, colCount int32
_ = binary.Read(f, binary.BigEndian, &rowCount)
_ = binary.Read(f, binary.BigEndian, &colCount)
if rowCount != 28 || colCount != 28 {
return nil, fmt.Errorf("invalid number of rows or columns in the file: %v", filename)
}
imageSize := rowCount * colCount
images := make([][]float32, 0, imageCount)
for range imageCount {
pixels := make([]float32, 0, imageSize)
for range imageSize {
var pixel uint8
err = binary.Read(f, binary.BigEndian, &pixel)
if err != nil {
return nil, fmt.Errorf("invalid pixel in the file: %v, %w", filename, err)
}
normalized := float32(pixel) / 255.0
pixels = append(pixels, normalized)
}
images = append(images, pixels)
}
return images, nil
}