- refactored

- use Armadillo for load/save of weight and training data

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@775 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-10 15:25:52 +00:00
parent 952cf26930
commit ff2086a1ff
9 changed files with 125 additions and 65 deletions
+45 -23
View File
@@ -175,7 +175,7 @@ bool Stack::loadWeights()
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (!pLayer->loadWeights(m_dir + "/" + m_name))
if (!pLayer->weightsLoad(m_dir, m_name))
{
return false;
}
@@ -189,7 +189,7 @@ bool Stack::saveWeights()
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (!pLayer->saveWeights(m_dir + "/" + m_name))
if (!pLayer->weightsSave(m_dir, m_name))
{
return false;
}
@@ -204,20 +204,20 @@ void Stack::train(Rbm::IListener* pListener)
while(pLayer)
{
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
pLayer->setBatch(m_trainingData);
pLayer->setBatch(m_trainingBatch);
pLayer->train(pListener);
pLayer = pLayer->next;
}
}
arma::mat& Stack::trainingData()
arma::mat& Stack::trainingBatch()
{
return m_trainingData;
return m_trainingBatch;
}
arma::mat Stack::trainingData(Layer* pThatLayer)
arma::mat Stack::trainingBatch(Layer* pThatLayer)
{
arma::mat thisBatch = m_trainingData;
arma::mat thisBatch = m_trainingBatch;
Layer *pLayer = m_pLayers;
while (pLayer)
{
@@ -231,8 +231,19 @@ arma::mat Stack::trainingData(Layer* pThatLayer)
return thisBatch;
}
size_t Stack::loadTraining(bool doNormalize)
size_t Stack::loadTrainingBatch(bool doNormalize)
{
{
std::string path = m_dir + "/" + m_name + ".training.mat";
bool success = m_trainingBatch.load(path, arma::arma_ascii);
if (success)
{
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n";
return m_trainingBatch.n_rows;
}
}
uint32_t numTraining = 0;
uint32_t numVisible = 0;
@@ -255,8 +266,8 @@ size_t Stack::loadTraining(bool doNormalize)
{
return 0;
}
m_trainingData = arma::zeros(numTraining, numVisible);
m_trainingBatch = arma::zeros(numTraining, numVisible);
uint32_t i, j;
for (i=0; i < numTraining; i++)
{
@@ -266,7 +277,7 @@ size_t Stack::loadTraining(bool doNormalize)
int result = fscanf(pFile, "%f", &v);
if (result > 0)
{
m_trainingData(i, j) = v;
m_trainingBatch(i, j) = v;
}
}
}
@@ -275,13 +286,24 @@ size_t Stack::loadTraining(bool doNormalize)
if (doNormalize)
{
m_trainingData = Rbm::normalize(m_trainingData);
m_trainingBatch = Rbm::normalize(m_trainingData);
}
return numTraining;
}
size_t Stack::saveTraining()
size_t Stack::saveTrainingBatch()
{
{
std::string path = m_dir + "/" + m_name + ".training.mat";
bool success = m_trainingBatch.save(path, arma::arma_ascii);
if (success)
{
std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
return m_trainingBatch.n_rows;
}
}
std::string filename = m_dir + "/" + m_name + ".training.dat";
FILE *pFile = fopen(filename.c_str(), "w");
@@ -291,37 +313,37 @@ size_t Stack::saveTraining()
return 0;
}
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_rows);
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_cols);
fprintf(pFile, "%u\n", (uint32_t)m_trainingBatch.n_rows);
fprintf(pFile, "%u\n", (uint32_t)m_trainingBatch.n_cols);
uint32_t i, j;
for (i=0; i < m_trainingData.n_rows; i++)
for (i=0; i < m_trainingBatch.n_rows; i++)
{
for (j=0; j < m_trainingData.n_cols; j++)
for (j=0; j < m_trainingBatch.n_cols; j++)
{
fprintf(pFile, "%3.6f\n", m_trainingData(i, j));
fprintf(pFile, "%3.6f\n", m_trainingBatch(i, j));
}
}
fclose(pFile);
std::cout << "Saved " << m_trainingData.n_rows << " training samples\n";
return m_trainingData.n_rows;
std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
return m_trainingBatch.n_rows;
}
size_t Stack::numTraining()
{
return m_trainingData.n_rows;
return m_trainingBatch.n_rows;
}
void Stack::addTraining(const arma::mat &toAdd)
{
m_trainingData.insert_rows(m_trainingData.n_rows, toAdd);
m_trainingBatch.insert_rows(m_trainingBatch.n_rows, toAdd);
}
void Stack::delTraining(int index)
{
m_trainingData.shed_row(index);
m_trainingBatch.shed_row(index);
}