- 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:
2022-01-06 11:00:51 +00:00
parent c2834d9f23
commit cbadd106f4
+22 -23
View File
@@ -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