- Layer: removed legacy training load/save
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@779 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+4
-71
@@ -239,54 +239,13 @@ size_t Stack::loadTrainingBatch(bool doNormalize)
|
||||
|
||||
if (success)
|
||||
{
|
||||
if (doNormalize)
|
||||
{
|
||||
m_trainingBatch = Rbm::normalize(m_trainingBatch);
|
||||
}
|
||||
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n";
|
||||
return m_trainingBatch.n_rows;
|
||||
}
|
||||
|
||||
uint32_t numTraining = 0;
|
||||
uint32_t numVisible = 0;
|
||||
|
||||
FILE *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;
|
||||
}
|
||||
m_trainingBatch = 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)
|
||||
{
|
||||
m_trainingBatch(i, j) = v;
|
||||
}
|
||||
}
|
||||
}
|
||||
fclose(pFile);
|
||||
std::cout << "Loaded " << numTraining << " training samples\n";
|
||||
|
||||
if (doNormalize)
|
||||
{
|
||||
m_trainingBatch = Rbm::normalize(m_trainingBatch);
|
||||
}
|
||||
return numTraining;
|
||||
}
|
||||
|
||||
size_t Stack::saveTrainingBatch()
|
||||
@@ -299,32 +258,6 @@ size_t Stack::saveTrainingBatch()
|
||||
std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
|
||||
return m_trainingBatch.n_rows;
|
||||
}
|
||||
|
||||
FILE *pFile = fopen(filename.c_str(), "w");
|
||||
|
||||
if (!pFile)
|
||||
{
|
||||
std::cout << "Could not open " << filename << "!" << std::endl;
|
||||
return 0;
|
||||
}
|
||||
|
||||
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_trainingBatch.n_rows; i++)
|
||||
{
|
||||
for (j=0; j < m_trainingBatch.n_cols; j++)
|
||||
{
|
||||
fprintf(pFile, "%3.6f\n", m_trainingBatch(i, j));
|
||||
}
|
||||
}
|
||||
|
||||
fclose(pFile);
|
||||
|
||||
std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
|
||||
return m_trainingBatch.n_rows;
|
||||
}
|
||||
|
||||
size_t Stack::numTraining()
|
||||
|
||||
Reference in New Issue
Block a user