- moved context awareness from Rbm to Layer
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@769 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+18
-2
@@ -37,8 +37,24 @@ public:
|
|||||||
void train(arma::mat const &batch, IListener *pListener=nullptr)
|
void train(arma::mat const &batch, IListener *pListener=nullptr)
|
||||||
{
|
{
|
||||||
if (batch.n_rows > 0)
|
if (batch.n_rows > 0)
|
||||||
{
|
{
|
||||||
Rbm::train(trainingData(batch), pListener);
|
if (numContext())
|
||||||
|
{
|
||||||
|
arma::mat c_states = arma::zeros(batch.n_rows, numContext());
|
||||||
|
arma::mat ctx = arma::zeros(1, numContext());
|
||||||
|
for (int i=1; i < batch.n_rows; i++)
|
||||||
|
{
|
||||||
|
arma::mat h = prob(v_to_h(arma::join_rows(batch.row(i-1), ctx)));
|
||||||
|
ctx = h;
|
||||||
|
c_states.row(i) = ctx;
|
||||||
|
}
|
||||||
|
arma::mat batch_with_ctx = arma::join_rows(batch, c_states);
|
||||||
|
Rbm::train(trainingData(batch_with_ctx), pListener);
|
||||||
|
}
|
||||||
|
else
|
||||||
|
{
|
||||||
|
Rbm::train(trainingData(batch), pListener);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+1
-20
@@ -147,30 +147,13 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
|||||||
while (trainingSizeRemain && !shouldAbort)
|
while (trainingSizeRemain && !shouldAbort)
|
||||||
{
|
{
|
||||||
int miniBatchSizeActual = std::min(m_params.miniBatchSize, trainingSizeRemain);
|
int miniBatchSizeActual = std::min(m_params.miniBatchSize, trainingSizeRemain);
|
||||||
arma::mat miniBatch_v = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
||||||
trainingSizeRemain -= miniBatchSizeActual;
|
trainingSizeRemain -= miniBatchSizeActual;
|
||||||
batchRowIndex += miniBatchSizeActual;
|
batchRowIndex += miniBatchSizeActual;
|
||||||
int scaler = std::min(m_params.miniBatchSize, (int)batch.n_rows);
|
int scaler = std::min(m_params.miniBatchSize, (int)batch.n_rows);
|
||||||
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 miniBatch;
|
|
||||||
if (numContext() > 0)
|
|
||||||
{
|
|
||||||
arma::mat c_states(arma::zeros(miniBatchSizeActual, numContext()));
|
|
||||||
|
|
||||||
for (int i=1; i < miniBatchSizeActual; i++)
|
|
||||||
{
|
|
||||||
arma::mat h = prob(v_to_h(arma::join_rows(miniBatch_v.row(i-1), ctx)));
|
|
||||||
ctx = h;
|
|
||||||
c_states.row(i) = ctx;
|
|
||||||
}
|
|
||||||
miniBatch = (arma::join_rows(miniBatch_v, c_states));
|
|
||||||
}
|
|
||||||
else
|
|
||||||
{
|
|
||||||
miniBatch = miniBatch_v;
|
|
||||||
}
|
|
||||||
arma::mat v_states(miniBatch);
|
arma::mat v_states(miniBatch);
|
||||||
|
|
||||||
// Create hidden layer base on training data
|
// Create hidden layer base on training data
|
||||||
@@ -221,7 +204,6 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
|||||||
|
|
||||||
} // number of mini batches
|
} // number of mini batches
|
||||||
|
|
||||||
#if FIXED_BATCH
|
|
||||||
arma::mat diffErr = batch - prob(h_to_v(prob(v_to_h(batch))));
|
arma::mat diffErr = batch - prob(h_to_v(prob(v_to_h(batch))));
|
||||||
arma::mat diffErr_squared = diffErr % diffErr;
|
arma::mat diffErr_squared = diffErr % diffErr;
|
||||||
status.err_total = accu(diffErr_squared)/diffErr_squared.n_elem;
|
status.err_total = accu(diffErr_squared)/diffErr_squared.n_elem;
|
||||||
@@ -230,7 +212,6 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
|||||||
{
|
{
|
||||||
pListener->onProgress(this, status);
|
pListener->onProgress(this, status);
|
||||||
}
|
}
|
||||||
#endif
|
|
||||||
}
|
}
|
||||||
|
|
||||||
arma::mat Rbm::prob(const arma::mat &src)
|
arma::mat Rbm::prob(const arma::mat &src)
|
||||||
|
|||||||
+1
-1
@@ -144,7 +144,7 @@ public:
|
|||||||
|
|
||||||
arma::mat toHiddenProbs(const arma::mat &visible) const
|
arma::mat toHiddenProbs(const arma::mat &visible) const
|
||||||
{
|
{
|
||||||
return Rbm::prob(v_to_h(arma::join_rows(visible, m_ctx)));
|
return Rbm::prob(v_to_h(visible));
|
||||||
}
|
}
|
||||||
|
|
||||||
arma::mat toVisibleProbs(const arma::mat &hidden) const
|
arma::mat toVisibleProbs(const arma::mat &hidden) const
|
||||||
|
|||||||
Reference in New Issue
Block a user