- 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:
2022-01-10 08:09:49 +00:00
parent 967df2906d
commit e9de02ac3d
3 changed files with 20 additions and 23 deletions
+18 -2
View File
@@ -37,8 +37,24 @@ public:
void train(arma::mat const &batch, IListener *pListener=nullptr)
{
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
View File
@@ -147,30 +147,13 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
while (trainingSizeRemain && !shouldAbort)
{
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;
batchRowIndex += miniBatchSizeActual;
int scaler = std::min(m_params.miniBatchSize, (int)batch.n_rows);
double learning_rate = m_params.learningRate/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);
// Create hidden layer base on training data
@@ -221,7 +204,6 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
} // number of mini batches
#if FIXED_BATCH
arma::mat diffErr = batch - prob(h_to_v(prob(v_to_h(batch))));
arma::mat diffErr_squared = diffErr % diffErr;
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);
}
#endif
}
arma::mat Rbm::prob(const arma::mat &src)
+1 -1
View File
@@ -144,7 +144,7 @@ public:
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