Repository navigation
Expand file tree
/
Copy pathmain.cpp
More file actions
103 lines (78 loc) · 3 KB
/
Copy pathmain.cpp
File metadata and controls
103 lines (78 loc) · 3 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
#include "matrix.hpp"
#include "mnist.hpp"
#include "serialize.hpp"
#include "value.hpp"
#include <chrono>
#include <memory>
int main() {
auto start = std::chrono::high_resolution_clock::now();
std::vector<Matrix> train_images =
load_images("data/train-images-idx3-ubyte");
std::vector<Matrix> train_labels =
load_labels("data/train-labels-idx1-ubyte");
Matrix W1(128, 784, true);
Matrix B1(128, 1, false);
Matrix W2(64, 128, true);
Matrix B2(64, 1, false);
Matrix W3(10, 64, true);
Matrix B3(10, 1, false);
auto w1 = std::make_shared<Value>(W1);
auto b1 = std::make_shared<Value>(B1);
auto w2 = std::make_shared<Value>(W2);
auto b2 = std::make_shared<Value>(B2);
auto w3 = std::make_shared<Value>(W3);
auto b3 = std::make_shared<Value>(B3);
std::vector<std::shared_ptr<Value>> params = {w1, b1, w2, b2, w3, b3};
Scalar lr = 0.01;
i64 epochs = 10;
size_t N = train_images.size();
for (i64 epoch = 0; epoch < epochs; ++epoch) {
Scalar total_loss = 0.0;
i64 correct = 0;
for (size_t ex = 0; ex < N; ++ex) {
Matrix current_image = train_images[ex];
Matrix current_label = train_labels[ex];
auto current_img = std::make_shared<Value>(current_image);
for (auto &p : params) {
p->grad = Matrix(p->grad.rows, p->grad.cols, false);
}
auto h1 = tanh_(add(matmul(w1, current_img), b1));
auto h2 = tanh_(add(matmul(w2, h1), b2));
auto logits = add(matmul(w3, h2), b3);
auto loss = cross_entropy(logits, current_label);
loss->backward();
for (auto p : params) {
p->data = p->data - p->grad.scalar_multiply(lr);
}
total_loss += loss->data.at(0, 0);
if (argmax(logits->data) == argmax(current_label))
correct++;
}
std::cout << "epoch: " << epoch << " loss: " << total_loss / N
<< " acc: " << correct * 100.0 / N << "%\n";
}
save_matrix(w1->data, "model/w1.txt");
save_matrix(b1->data, "model/b1.txt");
save_matrix(w2->data, "model/w2.txt");
save_matrix(b2->data, "model/b2.txt");
save_matrix(w3->data, "model/w3.txt");
save_matrix(b3->data, "model/b3.txt");
auto end = std::chrono::high_resolution_clock::now();
auto seconds =
std::chrono::duration_cast<std::chrono::seconds>(end - start).count();
std::cout << "training took " << seconds << " seconds\n";
std::vector<Matrix> test_images = load_images("data/t10k-images-idx3-ubyte");
std::vector<Matrix> test_labels = load_labels("data/t10k-labels-idx1-ubyte");
i64 test_correct = 0;
for (size_t ex = 0; ex < test_images.size(); ++ex) {
auto x = std::make_shared<Value>(test_images[ex]);
auto h1 = tanh_(add(matmul(w1, x), b1));
auto h2 = tanh_(add(matmul(w2, h1), b2));
auto o = add(matmul(w3, h2), b3); // o is a Value
if (argmax(o->data) == argmax(test_labels[ex])) // read ->data for argmax
test_correct++;
}
Scalar test_acc = (Scalar)test_correct / test_images.size() * 100;
std::cout << "TEST accuracy: " << test_acc << "%\n";
return 0;
}