- 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:
+45
-23
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user