- speed improvment: sample minibatch only on minibatch change
- refactored git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@744 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+22
-23
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user