Num Rec app and save & loading

This commit is contained in:
IlyaShurupov 2023-10-26 16:27:43 +03:00
parent fc20f3594d
commit b7a89b714e
12 changed files with 8171 additions and 12 deletions

View file

@ -0,0 +1,77 @@
#include "FCNN.hpp"
#include "LocalConnection.hpp"
// #include "NewPlacement.hpp"
#define STB_IMAGE_WRITE_IMPLEMENTATION
#define STB_IMAGE_IMPLEMENTATION
#include "stb_image.h"
#include "stb_image_write.h"
using namespace tp;
void loadImage(Buffer<halnf>& output, const char* name) {
int x, y, channels_in_file;
unsigned char* loadedImage = stbi_load(name, &x, &y, &channels_in_file, 4);
if (!loadedImage) return;
output.reserve(x * y);
for (auto i : Range(output.size())) {
output[i] = loadedImage[i * 4] / 255.f;
}
stbi_image_free(loadedImage);
}
void loadNN(FCNN& nn) {
ArchiverLocalConnection<true> archiver;
archiver.connection.connect(LocalConnection::Location("NumRec.wb"), LocalConnection::Type(true));
if (archiver.connection.getConnectionStatus().isOpened()) {
archiver % nn;
} else {
Buffer<halni> layers = { 784, 10 };
nn.initializeRandom(layers);
}
}
void executeCmd(const char* imageName) {
FCNN nn;
Buffer<halnf> output(10);
Buffer<halnf> input;
loadNN(nn);
loadImage(input, imageName);
nn.evaluate(input, output);
printf("Output - ");
for (auto val : output) {
printf("%f ", val.data());
}
printf("\n\n");
}
int main(int argc, char** argv) {
const char* imageName = "digit.png";
if (argc == 2) {
imageName = argv[1];
}
ModuleManifest* deps[] = { &gModuleDataAnalysis, &gModuleConnection, nullptr };
ModuleManifest module = ModuleManifest("NumRec", nullptr, nullptr, deps);
if (!module.initialize()) {
return 1;
}
executeCmd(imageName);
module.deinitialize();
return 0;
}

View file

@ -1,9 +1,27 @@
#include "FCNN.hpp"
#include "LocalConnection.hpp"
#include "NewPlacement.hpp"
#define STB_IMAGE_WRITE_IMPLEMENTATION
#include "stb_image_write.h"
using namespace tp;
void writeImage(const Buffer<halnf>& image, const char* name) {
struct Tmp {
uint1 r, g, b, a;
};
Buffer<Tmp> converted;
converted.reserve(image.size());
for (auto i = 0; i < image.size(); i++) {
auto val = uint1(image[i] * 255);
converted[i] = { val, val, val, 255 };
}
stbi_write_png(name, 28, 28, 4, converted.getBuff(), 28 * 4);
}
struct Dataset {
ualni length = 0;
Pair<ualni, ualni> imageSize = { 0, 0 };
@ -48,8 +66,20 @@ bool loadDataset(Dataset& out, const String& location) {
struct NumberRec {
NumberRec() {
Buffer<halni> layers = { 784, 10 };
nn.initializeRandom(layers);
// try to load wb file
{
ArchiverLocalConnection<true> archiver;
archiver.connection.connect(LocalConnection::Location("NumRec.wb"), LocalConnection::Type(true));
if (archiver.connection.getConnectionStatus().isOpened()) {
archiver % nn;
} else {
Buffer<halni> layers = { 784, 10 };
nn.initializeRandom(layers);
}
}
Dataset dataset;
@ -79,6 +109,20 @@ struct NumberRec {
}
output.reserve(10);
writeImage(mTestcases.first().input, "tmp1.png");
writeImage(mTestcases.last().input, "tmp2.png");
}
~NumberRec() {
// save aas wb file
{
ArchiverLocalConnection<false> archiver;
archiver.connection.connect(LocalConnection::Location("NumRec.wb"), LocalConnection::Type(false));
if (archiver.connection.getConnectionStatus().isOpened()) {
archiver % nn;
}
}
}
halnf eval(ualni idx) {
@ -179,7 +223,7 @@ int main() {
auto batchSize = trainRange.idxDiff() / numBatches;
for (auto epoch : Range(10)) {
for (auto epoch : Range(1)) {
printf("Epoch %i\n", epoch.index());
for (auto batchIdx : Range(trainRange.idxDiff() / batchSize)) {
@ -204,7 +248,7 @@ int main() {
// app.displayImage(i);
}
printf("\n\nIncorrect - %i out of %i (%f)\n\n", errors, testRange.idxDiff(), (halnf) errors / (halnf) testRange.idxDiff() );
printf("\n\nIncorrect - %i out of %i (%f)\n\n", errors, testRange.idxDiff(), (halnf) errors / (halnf) testRange.idxDiff());
}
module.deinitialize();

Binary file not shown.