#include #include #include #include #include #include #include #include #include "Rbm.hpp" #include "Layer.hpp" using namespace std; class RbmListener : public Rbm::IListener { public: RbmListener() {} virtual ~RbmListener() {} bool onProgress(const Rbm::Status &status) { std::cout << "Progress : " << 100*status.progress << " %" << std::endl; std::cout << "epoch : " << status.epoch << std::endl; std::cout << "trainingSizeRemain: " << status.trainingSizeRemain << std::endl; std::cout << "error (per mini batch) = " << status.err << std::endl; std::cout << "error (total) = " << status.err_total << std::endl; std::cout << "L1 = " << status.L1 << std::endl; std::cout << "L2 = " << status.L2 << std::endl; return true; } }; arma::mat loadTraining(const string &filename) { uint32_t numTraining = 0; uint32_t numVisible = 0; FILE *pFile; pFile = fopen(filename.c_str(), "r"); if (!pFile) { std::cout << "Could not open " << filename << "!" << std::endl; return 0; } int result = fscanf(pFile, "%d\n", &numTraining); if (result < 0) { return 0; } result = fscanf(pFile, "%d\n", &numVisible); if (result < 0) { return 0; } arma::mat data = arma::zeros(numTraining, numVisible); uint32_t i, j; for (i=0; i < numTraining; i++) { for (j=0; j < numVisible; j++) { float v; int result = fscanf(pFile, "%f", &v); if (result > 0) { data(i, j) = v; } } } fclose(pFile); return data; } void saveWeight(const string &filename, size_t numVisibleX, size_t numVisibleY, const arma::mat &w, const arma::mat &bv, const arma::mat &bh) { FILE *pFile; pFile = fopen(filename.c_str(), "w"); if (!pFile) { std::cout << "Could not open " << filename << "!" << std::endl; return; } size_t numHidden = bh.n_elem; size_t numVisible = bv.n_elem; fprintf(pFile, "%d %d %d\n", (int)numVisibleX, (int)numVisibleY, (int)numHidden); uint32_t i, j; for (i=0; i < numVisible; i++) { fprintf(pFile, "%3.6f\n", bv(i)); } for (i=0; i < numHidden; i++) { fprintf(pFile, "%3.6f\n", bh(i)); } for (i=0; i < numVisible; i++) { for (j=0; j < numHidden; j++) { fprintf(pFile, "%3.6f ", w(i,j)); } fprintf(pFile, "\n"); } fclose(pFile); } void saveProject(const string &prjname, const Layer &layer, size_t numTraining) { ofstream ofs(prjname + string(".prj")); Json::StyledWriter writer; Json::Value project; project["project"]["name"] = prjname; project["project"]["num_training"] = (int)numTraining; project["project"]["training_file"] = prjname + string(".training.dat"); Json::Value jsonLayer = layer.toJson(); Json::Value jsonLayers(Json::arrayValue); jsonLayers.append(jsonLayer); project["project"]["layers"] = jsonLayers; ofs << writer.write(project); } int main() { printf("Hallo, Welt!\n"); const string project("mnist"); RbmListener statusDisplay; Rbm::Params rbmParams; arma::mat batch = loadTraining(project + string(".training.dat")); size_t numTraining = batch.n_rows; size_t numVisible = batch.n_cols; size_t numHidden = 64; Layer layer(project, 0, numVisible, numHidden); saveProject(project, layer, numTraining); printf("Loaded %d training samples\n", (int)numTraining); arma::mat w = arma::zeros(numVisible, numHidden); arma::mat bv = arma::zeros(1, numVisible); arma::mat bh = arma::zeros(1, numHidden); Rbm rbm(rbmParams, numVisible, numHidden); rbm.train(batch, 1000, 100, &statusDisplay); saveWeight(project + string(".weights.dat"), 28, 28, rbm.w(), rbm.bv(), rbm.bh()); arma::mat v = arma::randu(numTraining, numVisible); arma::mat h = rbm.toHiddenProbs(v); arma::mat r = rbm.toVisibleProbs(h); return 0; }