-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.cpp
More file actions
123 lines (105 loc) · 4.33 KB
/
Copy pathmain.cpp
File metadata and controls
123 lines (105 loc) · 4.33 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
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
#include <chrono> // for timing
#include <thread> // for sleeping
#include "src/Graphics/Graphics.h"
#include "src/NeuralNetwork/NeuralNetwork.h"
#include "src/MNISTLoader/MNISTLoader.h"
#include "src/NeuralNetwork/CLProgram/CLProgram.h"
std::pair<std::vector<std::vector<float>>, std::vector<std::vector<float>>> convert(std::pair<std::vector<std::vector<uint8_t>>, std::vector<std::vector<float>>> in) {
std::vector<std::vector<float>> floatInput;
// Convert input values to floats between 0 and 1
for (const auto& row : in.first) {
std::vector<float> floatRow;
for (const auto& value : row) {
float floatValue = static_cast<float>(value) / 255.0f;
floatRow.push_back(floatValue);
}
floatInput.push_back(floatRow);
}
return std::make_pair(floatInput, in.second);
}
std::pair<std::vector<std::vector<float>>, std::vector<std::vector<float>>> loadMNISTData()
{
std::string input_path = "MNIST/";
std::string training_images_filepath = input_path + "train-images-idx3-ubyte/train-images-idx3-ubyte";
std::string training_labels_filepath = input_path + "train-labels-idx1-ubyte/train-labels-idx1-ubyte";
std::string test_images_filepath = input_path + "t10k-images-idx3-ubyte/t10k-images-idx3-ubyte";
std::string test_labels_filepath = input_path + "t10k-labels-idx1-ubyte/t10k-labels-idx1-ubyte";
MnistDataloader mnistData = MnistDataloader(training_images_filepath, training_labels_filepath, test_images_filepath, test_labels_filepath);
auto test = mnistData.load_data_f();
return convert(test);// mnistData.load_data_f();
}
int main() {
Graphics graphics = Graphics();
/*
std::vector<std::vector<float>> inputs = {
{ 0.0f, 0.0f },
{ 1.0f, 0.0f },
{ 0.0f, 1.0f },
{ 1.0f, 1.0f },
};
std::vector<std::vector<float>> outputs = {
{ 0.0f },
{ 1.0f },
{ 1.0f },
{ 1.0f }
};
NetworkParams perceptronNetworkParams = NetworkParams(
"src/kernels/perceptron.cl",
inputs[0].size(),
outputs[0].size(),
0,
{ { 0, { 4, 4, } }, },
inputs.size()
);
int batchSize = 4;
float learningRate = 1.0f;
NeuralNetwork<float> network = NeuralNetwork<float>(perceptronNetworkParams);
std::pair<std::vector<std::vector<float>>, std::vector<std::vector<float>>> trainingData = std::make_pair(inputs, outputs);
*/
//std::pair<std::vector<std::vector<float>>, std::vector<std::vector<float>>> trainingData = loadMNISTData();
std::pair<std::vector<std::vector<float>>, std::vector<std::vector<float>>> trainingData = {
{
{ 0.0f, 0.0f, },
{ 1.0f, 0.0f, },
{ 0.0f, 1.0f, },
{ 1.0f, 1.0f, },
},
{
{ 0.0f },
{ 1.0f },
{ 1.0f },
{ 1.0f }
}
};
NetworkParams mnistNetworkParams = NetworkParams(
"src/kernels/perceptron.cl",
trainingData.first[0].size(),
trainingData.second[0].size(),
0,
{ },
trainingData.first.size()
);
NeuralNetwork<float> network = NeuralNetwork<float>(mnistNetworkParams);
int batchSize = 4;
int learningRate = 1.0f;
GLuint iterations = 10;
network.train(trainingData, iterations, iterations / 10, batchSize, learningRate);
int iter = 0;
auto last_prediction_time = std::chrono::high_resolution_clock::now(); // initialize timer
// Main loop
while (graphics.is_running()) {
// Check if 0.5 seconds have passed since last prediction
auto now = std::chrono::high_resolution_clock::now();
if (!network.training && std::chrono::duration_cast<std::chrono::milliseconds>(now - last_prediction_time).count() >= 500) {
network.predict(trainingData.first[iter], trainingData.second[iter]);
iter = (iter + 1) % trainingData.second.size();
last_prediction_time = now; // update last prediction time
}
graphics.setupScene();
graphics.drawNeurons(network.returnNetworkValues(), network.returnWeightValues(), network.returnBiasValues());
graphics.swapBuffersAndPoll();
std::this_thread::sleep_for(std::chrono::milliseconds(16)); // limit frame rate to ~60 fps
}
// Clean up
return 0;
}