#include #include #include #include #include "Rbm.hpp" 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 char *pFilename) { uint32_t numTraining = 0; uint32_t numVisible = 0; FILE *pFile; pFile = fopen(pFilename,"r"); if (!pFile) { 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 char *pFilename, size_t numVisibleX, size_t numVisibleY, const arma::mat &w, const arma::mat &bv, const arma::mat &bh) { FILE *pFile; pFile = fopen(pFilename,"w"); if (!pFile) 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); } int main() { printf("Hallo, Welt!\n"); RbmListener statusDisplay; Rbm::Params params; arma::mat batch = loadTraining("mnist_2.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("mnist_2.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; }