- moved numEpochs and miniBatchSize to RBM::Params

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@606 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-06 06:12:51 +00:00
parent c834a86f0f
commit 3ef7b07f5e
6 changed files with 31 additions and 15 deletions
+3 -3
View File
@@ -57,7 +57,7 @@ Json::Value Rbm::toJson() const
return rbm;
}
void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener)
void Rbm::train(const arma::mat& batch, IListener* pListener)
{
Status status;
size_t epoch;
@@ -83,7 +83,7 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs,
while (status.trainingSizeRemain)
{
size_t miniBatchSizeActual = std::min(miniBatchSize, status.trainingSizeRemain);
size_t miniBatchSizeActual = std::min(m_params.miniBatchSize, status.trainingSizeRemain);
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
status.trainingSizeRemain -= miniBatchSizeActual;
batchRowIndex += miniBatchSizeActual;
@@ -95,7 +95,7 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs,
arma::mat hid_state(miniBatchSizeActual, m_w.n_cols);
arma::mat hid_probs(miniBatchSizeActual, m_w.n_cols);
for (epoch=0; epoch < numEpochs; epoch++)
for (epoch=0; epoch < m_params.numEpochs; epoch++)
{
// Create hidden layer base on training data