#include "FullyConnectedNN.hpp" #include "LocalConnection.hpp" #include "NewPlacement.hpp" using namespace tp; struct Dataset { ualni length = 0; Pair imageSize = { 0, 0 }; Buffer labels; Buffer> images; }; bool loadDataset(Dataset& out, const String& location) { LocalConnection dataset; dataset.connect(LocalConnection::Location(location), LocalConnection::Type(true)); if (!dataset.getConnectionStatus().isOpened()) { return false; } LocalConnection::Byte length; dataset.readBytes(&length, 1); LocalConnection::Byte sizeX; dataset.readBytes(&sizeX, 1); LocalConnection::Byte sizeY; dataset.readBytes(&sizeY, 1); out.length = ((ualni) length) * 1000; out.imageSize = { sizeX, sizeY }; out.labels.reserve(out.length); out.images.reserve(out.length); for (auto i : Range(out.length)) { auto& image = out.images[i]; image.reserve(sizeX * sizeY); dataset.readBytes((LocalConnection::Byte*) image.getBuff(), sizeX * sizeY); } LocalConnection::Byte label; dataset.readBytes((LocalConnection::Byte*) out.labels.getBuff(), out.length); return true; } halnf test(const Dataset& dataset, FullyConnectedNN& nn, Range range) { ualni numFailed = 0; for (auto i : range) { auto& image = dataset.images[i]; auto label = dataset.labels[i]; Buffer results; Buffer input; results.reserve(10); input.reserve(image.size()); for (auto pixelIdx : Range(image.size())) { input[pixelIdx] = (halnf) image[pixelIdx] / 255.f; } nn.evaluate(input, results); ualni resultNumber = 0; for (auto resIdx : Range(results.size())) { if (results[resIdx] > results[resultNumber]) { resultNumber = resIdx; } } if (resultNumber != label) { numFailed++; } } return (halnf) numFailed / (halnf) range.idxDiff(); } void numRec() { Dataset dataset; FullyConnectedNN nn; Buffer layers; layers = { 784, 128, 10 }; nn.initializeRandom(layers); if (!loadDataset(dataset, "rsc/mnist")) { printf("Cant Load Mnist Dataset\n"); return; } auto errorPercentage = test(dataset, nn, { 0, 100 }); printf("Percentage error : %f\n", errorPercentage); } int main() { ModuleManifest* deps[] = { &gModuleDataAnalysis, &gModuleConnection, nullptr }; ModuleManifest module = ModuleManifest("NumRec", nullptr, nullptr, deps); if (!module.initialize()) { return 1; } numRec(); module.deinitialize(); return 0; }