- Stack holds training data

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@636 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-07 19:30:25 +00:00
parent 4c68cd4370
commit 3d6a6fbf00
4 changed files with 47 additions and 43 deletions
+3 -3
View File
@@ -8,22 +8,22 @@
"numVisibleX" : 28, "numVisibleX" : 28,
"numVisibleY" : 28, "numVisibleY" : 28,
"rbm" : { "rbm" : {
"numHidden" : 256,
"numVisible" : 784,
"params" : { "params" : {
"doRaoBlackwell" : 1, "doRaoBlackwell" : 1,
"doSampleBatch" : 0, "doSampleBatch" : 0,
"gibbsDoSampleHidden" : 1, "gibbsDoSampleHidden" : 1,
"gibbsDoSampleVisible" : 0, "gibbsDoSampleVisible" : 0,
"learningRate" : 0.10000000000000001, "learningRate" : 0.10000000000000001,
"miniBatchSize" : 1000,
"momentum" : 0.5, "momentum" : 0.5,
"numEpochs" : 100,
"numGibbs" : 1, "numGibbs" : 1,
"weightDecay" : 0 "weightDecay" : 0
} }
}, },
"weights_file" : "Layer.0.weights.dat" "weights_file" : "Layer.0.weights.dat"
} }
], ],
"name" : "mnist" "name" : "mnist"
} }
} }
+1 -1
View File
@@ -628,7 +628,7 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
{ {
//[UserSliderCode_patterSlider] -- add your slider handling code here.. //[UserSliderCode_patterSlider] -- add your slider handling code here..
m_trainingIndex = (int)sliderThatWasMoved->getValue(); m_trainingIndex = (int)sliderThatWasMoved->getValue();
m_pLayer->setTrainingData(m_stack->training().row(m_trainingIndex)); m_pLayer->setTrainingData(m_stack->trainingData().row(m_trainingIndex));
//[/UserSliderCode_patterSlider] //[/UserSliderCode_patterSlider]
} }
else if (sliderThatWasMoved == WeightsSlider) else if (sliderThatWasMoved == WeightsSlider)
+33 -31
View File
@@ -182,38 +182,38 @@ bool Stack::saveWeights()
return true; return true;
} }
void Stack::train(const arma::mat& batch, Rbm::IListener* pListener) void Stack::train(Rbm::IListener* pListener)
{ {
Layer *pLayer = m_pLayers; Layer *pLayer = m_pLayers;
while(pLayer) while(pLayer)
{ {
train(pLayer->id(), batch, pListener); std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
pLayer->train(trainingData(pLayer), pListener);
pLayer = pLayer->next; pLayer = pLayer->next;
} }
} }
void Stack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) arma::mat& Stack::trainingData()
{ {
arma::mat thisBatch = batch; return m_trainingData;
Layer *pLayer = m_pLayers; }
while(pLayer)
{ arma::mat Stack::trainingData(Layer* pLayer)
if (pLayer->id() == layerId) {
arma::mat thisBatch = m_trainingData;
Layer *pThisLayer = m_pLayers;
while (pLayer) {
if (pThisLayer->id() == pLayer->id())
{ {
break; break;
} }
thisBatch = pLayer->toHiddenProbs(thisBatch); thisBatch = pLayer->toHiddenProbs(thisBatch);
pLayer = pLayer->next; pLayer = pLayer->next;
} }
return thisBatch;
if (pLayer)
{
std::cout << m_name << ": " << " Training of layer " << std::to_string(layerId) << std::endl;
pLayer->train(thisBatch, pListener);
}
} }
arma::mat Stack::loadTraining() size_t Stack::loadTraining()
{ {
uint32_t numTraining = 0; uint32_t numTraining = 0;
uint32_t numVisible = 0; uint32_t numVisible = 0;
@@ -237,7 +237,7 @@ arma::mat Stack::loadTraining()
{ {
return 0; return 0;
} }
arma::mat data = arma::zeros(numTraining, numVisible); m_trainingData = arma::zeros(numTraining, numVisible);
uint32_t i, j; uint32_t i, j;
for (i=0; i < numTraining; i++) for (i=0; i < numTraining; i++)
@@ -248,15 +248,15 @@ arma::mat Stack::loadTraining()
int result = fscanf(pFile, "%f", &v); int result = fscanf(pFile, "%f", &v);
if (result > 0) if (result > 0)
{ {
data(i, j) = v; m_trainingData(i, j) = v;
} }
} }
} }
fclose(pFile); fclose(pFile);
return data; return numTraining;
} }
void Stack::saveTraining(const arma::mat& batch) size_t Stack::saveTraining()
{ {
std::string filename = m_name + ".training.dat"; std::string filename = m_name + ".training.dat";
FILE *pFile = fopen(filename.c_str(), "w"); FILE *pFile = fopen(filename.c_str(), "w");
@@ -264,37 +264,39 @@ void Stack::saveTraining(const arma::mat& batch)
if (!pFile) if (!pFile)
{ {
std::cout << "Could not open " << filename << "!" << std::endl; std::cout << "Could not open " << filename << "!" << std::endl;
return; return 0;
} }
fprintf(pFile, "%u\n", (uint32_t)batch.n_rows); fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_rows);
fprintf(pFile, "%u\n", (uint32_t)batch.n_cols); fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_cols);
uint32_t i, j; uint32_t i, j;
for (i=0; i < batch.n_rows; i++) for (i=0; i < m_trainingData.n_rows; i++)
{ {
for (j=0; j < batch.n_cols; j++) for (j=0; j < m_trainingData.n_cols; j++)
{ {
fprintf(pFile, "%3.6f\n", batch(i, j)); fprintf(pFile, "%3.6f\n", m_trainingData(i, j));
} }
} }
fclose(pFile); fclose(pFile);
return m_trainingData.n_rows;
} }
size_t Stack::numTraining(const arma::mat &batch) size_t Stack::numTraining()
{ {
return batch.n_rows; return m_trainingData.n_rows;
} }
void Stack::addTraining(arma::mat &batch, const arma::mat &toAdd) void Stack::addTraining(const arma::mat &toAdd)
{ {
batch.insert_rows(batch.n_rows, toAdd); m_trainingData.insert_rows(m_trainingData.n_rows, toAdd);
} }
void Stack::delTraining(arma::mat &batch, int index) void Stack::delTraining(int index)
{ {
batch.shed_row(index); m_trainingData.shed_row(index);
} }
+10 -8
View File
@@ -44,23 +44,25 @@ public:
void delLayer(Layer *pLayer); void delLayer(Layer *pLayer);
Layer* getLayer(size_t layerId) const; Layer* getLayer(size_t layerId) const;
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener); void train(Rbm::IListener* pListener);
void train(const arma::mat& batch, Rbm::IListener* pListener);
bool load(LayerConstructor *pLayerConstructor=nullptr); bool load(LayerConstructor *pLayerConstructor=nullptr);
bool save(); bool save();
void weightsInit(double stddev); void weightsInit(double stddev);
bool loadWeights(); bool loadWeights();
bool saveWeights(); bool saveWeights();
static size_t numTraining(const arma::mat &batch); size_t numTraining();
static void addTraining(arma::mat &batch, const arma::mat &toAdd); void addTraining(const arma::mat &toAdd);
static void delTraining(arma::mat &batch, int index); void delTraining(int index);
arma::mat loadTraining(); size_t loadTraining();
void saveTraining(const arma::mat &batch); size_t saveTraining();
arma::mat& trainingData();
arma::mat trainingData(Layer *pLayer);
private: private:
std::string m_name; std::string m_name;
Layer *m_pLayers; Layer *m_pLayers;
arma::mat m_trainingData;
}; };