- 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);
}
+6 -6
View File
@@ -39,7 +39,9 @@ public:
Stack(const Stack& orig);
virtual ~Stack();
size_t numLayers();
void addLayer(Layer *pLayer);
void delLayer(Layer *pLayer);
Layer* getLayer(size_t layerId) const;
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener);
@@ -50,12 +52,10 @@ public:
bool loadWeights();
bool saveWeights();
void batchchanged()
{
std::cout << "Stack::batchchanged() called" << std::endl;
}
size_t numLayers();
static size_t numTraining(const arma::mat &batch);
static void addTraining(arma::mat &batch, const arma::mat &toAdd);
static void delTraining(arma::mat &batch, int index);
arma::mat loadTraining();
private:
std::string m_name;
+9 -4
View File
@@ -84,10 +84,15 @@ int main()
RbmListener statusDisplay;
Stack stack(project);
arma::mat batch = loadTraining(project + string(".training.dat"));
arma::mat batch = stack.loadTraining();
size_t numTraining = batch.n_rows;
printf("Loaded %d training samples\n", (int)numTraining);
printf("Loaded %d training samples\n", (int)batch.n_rows);
stack.addTraining(batch, batch.row(1));
printf("Loaded %d training samples\n", (int)batch.n_rows);
stack.delTraining(batch, 0);
printf("Loaded %d training samples\n", (int)batch.n_rows);
#if 1
const int numLayers = 4;
@@ -128,7 +133,7 @@ int main()
stack.saveWeights();
Layer *layer = stack.getLayer(0);
arma::mat v = arma::randu(numTraining, layer->bv().n_elem);
arma::mat v = arma::randu(batch.n_rows, layer->bv().n_elem);
arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h);
return 0;