- 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)
|
Rbm::Rbm(size_t numVisible, size_t numHidden)
|
||||||
: m_params()
|
: m_params()
|
||||||
, m_w(numVisible, numHidden)
|
, m_w(numVisible, numHidden)
|
||||||
, m_bv(1, numVisible)
|
|
||||||
, m_bh(1, numHidden)
|
, m_bh(1, numHidden)
|
||||||
|
, m_bv(1, numVisible)
|
||||||
{
|
{
|
||||||
assert(numVisible > 0);
|
assert(numVisible > 0);
|
||||||
assert(numHidden > 0);
|
assert(numHidden > 0);
|
||||||
@@ -28,8 +28,8 @@ Rbm::Rbm(size_t numVisible, size_t numHidden)
|
|||||||
Rbm::Rbm(const Rbm& orig)
|
Rbm::Rbm(const Rbm& orig)
|
||||||
: m_params(orig.m_params)
|
: m_params(orig.m_params)
|
||||||
, m_w(orig.m_w)
|
, m_w(orig.m_w)
|
||||||
, m_bv(orig.m_bv)
|
|
||||||
, m_bh(orig.m_bh)
|
, m_bh(orig.m_bh)
|
||||||
|
, m_bv(orig.m_bv)
|
||||||
{
|
{
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -41,9 +41,8 @@ Rbm::~Rbm()
|
|||||||
void Rbm::weightsInit(double stddev, double mu)
|
void Rbm::weightsInit(double stddev, double mu)
|
||||||
{
|
{
|
||||||
uniform(m_w, stddev, mu);
|
uniform(m_w, stddev, mu);
|
||||||
uniform(m_bv, stddev, mu);
|
|
||||||
uniform(m_bh, stddev, mu);
|
uniform(m_bh, stddev, mu);
|
||||||
|
uniform(m_bv, stddev, mu);
|
||||||
}
|
}
|
||||||
|
|
||||||
void Rbm::fromJson(Json::Value rbm)
|
void Rbm::fromJson(Json::Value rbm)
|
||||||
@@ -68,12 +67,12 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
|||||||
int lastProgress = -100;
|
int lastProgress = -100;
|
||||||
int batchRowIndex = 0;
|
int batchRowIndex = 0;
|
||||||
|
|
||||||
arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows));
|
arma::mat grad_bias_v(arma::zeros(1, m_bv.n_cols));
|
||||||
arma::mat grad_bias_h(arma::zeros(1, m_w.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 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_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_v(arma::zeros(1, m_bv.n_cols));
|
||||||
arma::mat momentum_bias_h(arma::zeros(1, m_w.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);
|
arma::mat penalty_weights = arma::zeros(m_w.n_rows, m_w.n_cols);
|
||||||
|
|
||||||
int trainingSizeRemain = batch.n_rows;
|
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 learning_rate = m_params.learningRate/scaler;
|
||||||
double weight_decay = m_params.weightDecay/scaler;
|
double weight_decay = m_params.weightDecay/scaler;
|
||||||
|
|
||||||
arma::mat vis_state(miniBatchSizeActual, m_w.n_rows);
|
arma::mat vis_state(miniBatchSizeActual, m_bv.n_cols);
|
||||||
arma::mat vis_probs(miniBatchSizeActual, m_w.n_rows);
|
arma::mat vis_probs(miniBatchSizeActual, m_bv.n_cols);
|
||||||
arma::mat hid_state(miniBatchSizeActual, m_w.n_cols);
|
arma::mat hid_state(miniBatchSizeActual, m_bh.n_cols);
|
||||||
arma::mat hid_probs(miniBatchSizeActual, m_w.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++)
|
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));
|
hid_probs = probsLogistic(toHiddenState(vis_state));
|
||||||
|
|
||||||
// Sample hidden
|
// Sample hidden
|
||||||
|
|||||||
Reference in New Issue
Block a user