#include #include #include #include #include #include #include #include #include "Rbm.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) { ofstream ofs(prjname + string(".prj")); Json::StyledWriter writer; Json::Value project; project["project"]["name"] = prjname; Json::Value layer1; layer1["name"] = "1"; layer1["num_hidden"] = 64; layer1["num_visible"] = 28*28; Json::Value params1; params1["weightInit"] = 0.01; params1["weightDecay"] = 0.001; params1["learningRate"] = 0.1; params1["momentum"] = 0.5; params1["doRaoBlackwell"] = 1; params1["gibbsDoSampleVisible"] = 0; params1["gibbsDoSampleHidden"] = 1; params1["doSampleBatch"] = 0; params1["numGibbs"] = 1; layer1["params"] = params1; Json::Value layer2; layer2["name"] = "2"; layer2["num_hidden"] = 16; layer2["num_visible"] = 64; Json::Value params2; params2["weightInit"] = 0.01; params2["weightDecay"] = 0.001; params2["learningRate"] = 0.1; params2["momentum"] = 0.5; params2["doRaoBlackwell"] = 1; params2["gibbsDoSampleVisible"] = 0; params2["gibbsDoSampleHidden"] = 1; params2["doSampleBatch"] = 0; params2["numGibbs"] = 1; layer2["params"] = params2; Json::Value layers(Json::arrayValue); layers.append(layer1); layers.append(layer2); project["project"]["layers"] = layers; ofs << writer.write(project); } int main() { printf("Hallo, Welt!\n"); const string project("mnist_2"); saveProject(project); RbmListener statusDisplay; Rbm::Params params; arma::mat batch = loadTraining(project + string(".training.dat")); size_t numTraining = batch.n_rows; size_t numVisible = batch.n_cols; size_t numHidden = 64; 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(params, w, bv, bh); rbm.train(batch, 1000, 100, &statusDisplay); saveWeight(project + string(".weights.dat"), 28, 28, w, bv, bh); arma::mat v = arma::randu(numTraining, numVisible); arma::mat h = rbm.toHiddenProbs(v); arma::mat r = rbm.toVisibleProbs(h); return 0; }