- added Stack::loadTraining(), Stack::addTraing(), Stack::delTraining()

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@623 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-07 15:24:23 +00:00
parent a5dd0acc58
commit 5d281ddca6
3 changed files with 79 additions and 10 deletions
+64
View File
@@ -61,6 +61,11 @@ void Stack::addLayer(Layer *pOtherLayer)
}
}
void Stack::delLayer(Layer* pLayer)
{
assert(!"Stack::delLayer: Not implemented!");
}
Layer* Stack::getLayer(size_t layerId) const
{
Layer *pLayer = m_pLayers;
@@ -208,3 +213,62 @@ void Stack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListe
}
}
arma::mat Stack::loadTraining()
{
uint32_t numTraining = 0;
uint32_t numVisible = 0;
FILE *pFile;
std::string filename = m_name + ".training.dat";
pFile = fopen(filename.c_str(), "r");
if (!pFile)
{
std::cout << "Could not open " << filename << "!" << std::endl;
return 0;
}
int result = fscanf(pFile, "%d\n", &numTraining);
if (result < 0)
{
return 0;
}
result = fscanf(pFile, "%d\n", &numVisible);
if (result < 0)
{
return 0;
}
arma::mat data = arma::zeros(numTraining, numVisible);
uint32_t i, j;
for (i=0; i < numTraining; i++)
{
for (j=0; j < numVisible; j++)
{
float v;
int result = fscanf(pFile, "%f", &v);
if (result > 0)
{
data(i, j) = v;
}
}
}
fclose(pFile);
return data;
}
size_t Stack::numTraining(const arma::mat &batch)
{
return batch.n_rows;
}
void Stack::addTraining(arma::mat &batch, const arma::mat &toAdd)
{
batch.insert_rows(batch.n_rows, toAdd);
}
void Stack::delTraining(arma::mat &batch, int index)
{
batch.shed_row(index);
}