diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 5eae51d..cfb919d 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -17,8 +17,8 @@ Rbm::Rbm(size_t numVisible, size_t numHidden) : m_params() , m_w(numVisible, numHidden) -, m_bv(1, numVisible) , m_bh(1, numHidden) +, m_bv(1, numVisible) { assert(numVisible > 0); assert(numHidden > 0); @@ -28,8 +28,8 @@ Rbm::Rbm(size_t numVisible, size_t numHidden) Rbm::Rbm(const Rbm& orig) : m_params(orig.m_params) , m_w(orig.m_w) -, m_bv(orig.m_bv) , m_bh(orig.m_bh) +, m_bv(orig.m_bv) { } @@ -41,9 +41,8 @@ Rbm::~Rbm() void Rbm::weightsInit(double stddev, double mu) { uniform(m_w, stddev, mu); - uniform(m_bv, stddev, mu); uniform(m_bh, stddev, mu); - + uniform(m_bv, stddev, mu); } void Rbm::fromJson(Json::Value rbm) @@ -68,12 +67,12 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) int lastProgress = -100; int batchRowIndex = 0; - arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows)); - arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols)); + arma::mat grad_bias_v(arma::zeros(1, m_bv.n_cols)); + arma::mat grad_bias_h(arma::zeros(1, m_bh.n_cols)); arma::mat grad_weight(arma::zeros(m_w.n_rows, m_w.n_cols)); arma::mat momentum_weights = arma::zeros(m_w.n_rows, m_w.n_cols); - arma::mat momentum_bias_v(arma::zeros(1, m_w.n_rows)); - arma::mat momentum_bias_h(arma::zeros(1, m_w.n_cols)); + arma::mat momentum_bias_v(arma::zeros(1, m_bv.n_cols)); + arma::mat momentum_bias_h(arma::zeros(1, m_bh.n_cols)); arma::mat penalty_weights = arma::zeros(m_w.n_rows, m_w.n_cols); int trainingSizeRemain = batch.n_rows; @@ -89,10 +88,21 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) double learning_rate = m_params.learningRate/scaler; double weight_decay = m_params.weightDecay/scaler; - arma::mat vis_state(miniBatchSizeActual, m_w.n_rows); - arma::mat vis_probs(miniBatchSizeActual, m_w.n_rows); - arma::mat hid_state(miniBatchSizeActual, m_w.n_cols); - arma::mat hid_probs(miniBatchSizeActual, m_w.n_cols); + arma::mat vis_state(miniBatchSizeActual, m_bv.n_cols); + arma::mat vis_probs(miniBatchSizeActual, m_bv.n_cols); + arma::mat hid_state(miniBatchSizeActual, m_bh.n_cols); + arma::mat hid_probs(miniBatchSizeActual, m_bh.n_cols); + + // Create hidden layer base on training data + if (m_params.doSampleBatch) + { + // When the hidden units are being driven by data, always use stochastic binary states + vis_state = sample(miniBatch); + } + else + { + vis_state = miniBatch; + } for (int epoch=0; epoch < m_params.numEpochs; epoch++) { @@ -110,17 +120,6 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) } } - // Create hidden layer base on training data - if (m_params.doSampleBatch) - { - // When the hidden units are being driven by data, always use stochastic binary states - vis_state = sample(miniBatch); - } - else - { - vis_state = miniBatch; - } - hid_probs = probsLogistic(toHiddenState(vis_state)); // Sample hidden