- Stack holds training data

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@636 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-07 19:30:25 +00:00
parent 4c68cd4370
commit 3d6a6fbf00
4 changed files with 47 additions and 43 deletions
+33 -31
View File
@@ -182,38 +182,38 @@ bool Stack::saveWeights()
return true;
}
void Stack::train(const arma::mat& batch, Rbm::IListener* pListener)
void Stack::train(Rbm::IListener* pListener)
{
Layer *pLayer = m_pLayers;
while(pLayer)
{
train(pLayer->id(), batch, pListener);
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
pLayer->train(trainingData(pLayer), pListener);
pLayer = pLayer->next;
}
}
void Stack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener)
arma::mat& Stack::trainingData()
{
arma::mat thisBatch = batch;
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (pLayer->id() == layerId)
return m_trainingData;
}
arma::mat Stack::trainingData(Layer* pLayer)
{
arma::mat thisBatch = m_trainingData;
Layer *pThisLayer = m_pLayers;
while (pLayer) {
if (pThisLayer->id() == pLayer->id())
{
break;
}
thisBatch = pLayer->toHiddenProbs(thisBatch);
pLayer = pLayer->next;
}
if (pLayer)
{
std::cout << m_name << ": " << " Training of layer " << std::to_string(layerId) << std::endl;
pLayer->train(thisBatch, pListener);
}
return thisBatch;
}
arma::mat Stack::loadTraining()
size_t Stack::loadTraining()
{
uint32_t numTraining = 0;
uint32_t numVisible = 0;
@@ -237,7 +237,7 @@ arma::mat Stack::loadTraining()
{
return 0;
}
arma::mat data = arma::zeros(numTraining, numVisible);
m_trainingData = arma::zeros(numTraining, numVisible);
uint32_t i, j;
for (i=0; i < numTraining; i++)
@@ -248,15 +248,15 @@ arma::mat Stack::loadTraining()
int result = fscanf(pFile, "%f", &v);
if (result > 0)
{
data(i, j) = v;
m_trainingData(i, j) = v;
}
}
}
fclose(pFile);
return data;
return numTraining;
}
void Stack::saveTraining(const arma::mat& batch)
size_t Stack::saveTraining()
{
std::string filename = m_name + ".training.dat";
FILE *pFile = fopen(filename.c_str(), "w");
@@ -264,37 +264,39 @@ void Stack::saveTraining(const arma::mat& batch)
if (!pFile)
{
std::cout << "Could not open " << filename << "!" << std::endl;
return;
return 0;
}
fprintf(pFile, "%u\n", (uint32_t)batch.n_rows);
fprintf(pFile, "%u\n", (uint32_t)batch.n_cols);
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_rows);
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_cols);
uint32_t i, j;
for (i=0; i < batch.n_rows; i++)
for (i=0; i < m_trainingData.n_rows; i++)
{
for (j=0; j < batch.n_cols; j++)
for (j=0; j < m_trainingData.n_cols; j++)
{
fprintf(pFile, "%3.6f\n", batch(i, j));
fprintf(pFile, "%3.6f\n", m_trainingData(i, j));
}
}
fclose(pFile);
return m_trainingData.n_rows;
}
size_t Stack::numTraining(const arma::mat &batch)
size_t Stack::numTraining()
{
return batch.n_rows;
return m_trainingData.n_rows;
}
void Stack::addTraining(arma::mat &batch, const arma::mat &toAdd)
void Stack::addTraining(const arma::mat &toAdd)
{
batch.insert_rows(batch.n_rows, toAdd);
m_trainingData.insert_rows(m_trainingData.n_rows, toAdd);
}
void Stack::delTraining(arma::mat &batch, int index)
void Stack::delTraining(int index)
{
batch.shed_row(index);
m_trainingData.shed_row(index);
}