- 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:
2022-01-10 15:25:52 +00:00
parent 952cf26930
commit ff2086a1ff
9 changed files with 125 additions and 65 deletions
+4 -13
View File
@@ -26,7 +26,6 @@ Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibl
, m_context(0, numContext) , m_context(0, numContext)
{ {
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl; cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
} }
Layer::Layer(const Layer& orig) Layer::Layer(const Layer& orig)
@@ -34,7 +33,6 @@ Layer::Layer(const Layer& orig)
, next(nullptr) , next(nullptr)
, prev(nullptr) , prev(nullptr)
, m_name(orig.m_name) , m_name(orig.m_name)
, m_weightsFile(orig.m_weightsFile)
, m_id(orig.m_id) , m_id(orig.m_id)
, m_numVisibleX(orig.m_numVisibleX) , m_numVisibleX(orig.m_numVisibleX)
, m_numVisibleY(orig.m_numVisibleY) , m_numVisibleY(orig.m_numVisibleY)
@@ -45,13 +43,10 @@ Layer::~Layer()
{ {
} }
bool Layer::loadWeights(const string &prjname) bool Layer::loadWeights(const string &prjname)
{ {
string filename = m_weightsFile; string filename = filePrefix(prjname) + ".weights.dat";
if (prjname.size() > 0)
{
filename = prjname + "." + m_weightsFile;
}
FILE *pFile = fopen(filename.c_str(),"r"); FILE *pFile = fopen(filename.c_str(),"r");
if (!pFile) if (!pFile)
@@ -108,11 +103,7 @@ bool Layer::saveWeights(const string &prjname)
{ {
int numHidden = m_bhv.n_elem; int numHidden = m_bhv.n_elem;
int numVisible = m_bv.n_elem; int numVisible = m_bv.n_elem;
string filename = m_weightsFile; string filename = filePrefix(prjname) + ".weights.dat";
if (prjname.size() > 0)
{
filename = prjname + "." + m_weightsFile;
}
FILE *pFile = fopen(filename.c_str(),"w"); FILE *pFile = fopen(filename.c_str(),"w");
if (!pFile) if (!pFile)
@@ -153,7 +144,7 @@ Json::Value Layer::toJson() const
Json::Value layer; Json::Value layer;
layer["name"] = m_name; layer["name"] = m_name;
layer["id"] = (int)m_id; layer["id"] = (int)m_id;
layer["weights_file"] = m_weightsFile; layer["weights_file"] = filePrefix("") + ".weights.dat";
layer["numVisibleX"] = (int)m_numVisibleX; layer["numVisibleX"] = (int)m_numVisibleX;
layer["numVisibleY"] = (int)m_numVisibleY; layer["numVisibleY"] = (int)m_numVisibleY;
layer["numHidden"] = (int)whv().n_cols; layer["numHidden"] = (int)whv().n_cols;
+50 -3
View File
@@ -31,8 +31,6 @@ public:
virtual ~Layer(); virtual ~Layer();
Json::Value toJson() const; Json::Value toJson() const;
bool loadWeights(const std::string &prjname="");
bool saveWeights(const std::string &prjname="");
void setBatch(arma::mat const &batch) void setBatch(arma::mat const &batch)
{ {
@@ -83,6 +81,43 @@ public:
return whv().n_cols; return whv().n_cols;
} }
bool weightsLoad(std::string const &dir, std::string const &prj)
{
arma::mat w;
arma::mat bh;
arma::mat bv;
bool result = true;
result &= w.load(filePrefix(prj) + ".w.dat", arma::arma_ascii);
result &= bh.load(filePrefix(prj) + ".bh.dat", arma::arma_ascii);
result &= bv.load(filePrefix(prj) + ".bv.dat", arma::arma_ascii);
if (result)
{
std::cout << "Layer " << m_id << ": Importing weights" << std::endl;
weightsAssign(w, bh, bv);
}
else
{
return loadWeights(prj);
}
return true;
}
bool weightsSave(std::string const &dir, std::string const &prj)
{
bool result = true;
result &= whv().save(filePrefix(prj) + ".w.dat", arma::arma_ascii);
result &= bh().save(filePrefix(prj) + ".bh.dat", arma::arma_ascii);
result &= bv().save(filePrefix(prj) + ".bv.dat", arma::arma_ascii);
if (result)
{
std::cout << "Layer " << m_id << ": Exporting weights" << std::endl;
}
return result;
}
arma::mat trainingData(arma::mat const &batch) arma::mat trainingData(arma::mat const &batch)
{ {
arma::mat thisBatch = batch; arma::mat thisBatch = batch;
@@ -131,13 +166,25 @@ public:
private: private:
std::string m_name; std::string m_name;
std::string m_weightsFile;
size_t m_id; size_t m_id;
size_t m_numVisibleX; size_t m_numVisibleX;
size_t m_numVisibleY; size_t m_numVisibleY;
size_t m_numContext; size_t m_numContext;
arma::mat m_context; arma::mat m_context;
// Compatibility
bool loadWeights(const std::string &prjname="");
bool saveWeights(const std::string &prjname="");
std::string filePrefix(const std::string &prjname) const
{
std::string filename = m_name + "." + std::to_string((int)m_id);
if (prjname.size() > 0)
{
filename = prjname + "." + filename;
}
return filename;
}
}; };
#endif /* RBMLAYER_HPP */ #endif /* RBMLAYER_HPP */
+4 -4
View File
@@ -814,7 +814,7 @@ void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged)
m_pLayer = static_cast<RbmComponent*>(m_stack->getLayer(index)); m_pLayer = static_cast<RbmComponent*>(m_stack->getLayer(index));
updateControls(); updateControls();
m_pLayer->redrawWeights(m_weightIndex); m_pLayer->redrawWeights(m_weightIndex);
if (m_stack->trainingData().n_rows > 0) if (m_stack->trainingBatch().n_rows > 0)
{ {
m_pLayer->setTrainingData(trainingAt(m_trainingIndex)); m_pLayer->setTrainingData(trainingAt(m_trainingIndex));
} }
@@ -878,7 +878,7 @@ void MainComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails
//[MiscUserCode] You can add your own definitions of your custom methods or any other code here... //[MiscUserCode] You can add your own definitions of your custom methods or any other code here...
void MainComponent::save () void MainComponent::save ()
{ {
m_stack->saveTraining(); m_stack->saveTrainingBatch();
m_stack->saveWeights(); m_stack->saveWeights();
m_stack->save(); m_stack->save();
} }
@@ -899,7 +899,7 @@ const juce::String& MainComponent::getBaseDir()
void MainComponent::run() void MainComponent::run()
{ {
trainButton->setButtonText (TRANS("Stop")); trainButton->setButtonText (TRANS("Stop"));
m_pLayer->setBatch(m_stack->trainingData()); m_pLayer->setBatch(m_stack->trainingBatch());
m_pLayer->train(this); m_pLayer->train(this);
trainButton->setButtonText (TRANS("Train")); trainButton->setButtonText (TRANS("Train"));
} }
@@ -916,7 +916,7 @@ bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
} }
RbmComponent *pComp = static_cast<RbmComponent*>(pRbm); RbmComponent *pComp = static_cast<RbmComponent*>(pRbm);
pComp->upPass(pComp->getTraining()); pComp->upPass(pComp->getTrainingPattern());
pComp->redrawReconstruction(); pComp->redrawReconstruction();
pComp->redrawWeights(); pComp->redrawWeights();
+5 -5
View File
@@ -98,12 +98,12 @@ private:
void clearTraining() void clearTraining()
{ {
patterSlider->setRange(0, 0, 1); patterSlider->setRange(0, 0, 1);
m_stack->trainingData().clear(); m_stack->trainingBatch().clear();
} }
void loadTraining() void loadTraining()
{ {
m_stack->loadTraining(rbmNormalizeDataToggleButton->getToggleState()); m_stack->loadTrainingBatch(rbmNormalizeDataToggleButton->getToggleState());
patterSlider->setRange(0, m_stack->numTraining()-1, 1); patterSlider->setRange(0, m_stack->numTraining()-1, 1);
} }
@@ -128,16 +128,16 @@ private:
{ {
if (m_pLayer->context().is_empty()) if (m_pLayer->context().is_empty())
{ {
m_pLayer->setBatch(m_stack->trainingData()); m_pLayer->setBatch(m_stack->trainingBatch());
} }
if (!m_pLayer->context().is_empty()) if (!m_pLayer->context().is_empty())
{ {
return arma::join_rows(m_stack->trainingData().row(index), m_pLayer->context().row(index)); return arma::join_rows(m_stack->trainingBatch().row(index), m_pLayer->context().row(index));
} }
else else
{ {
return m_stack->trainingData().row(index); return m_stack->trainingBatch().row(index);
} }
} }
+5 -5
View File
@@ -215,11 +215,11 @@ void RbmComponent::onDraw(DrawComponent &obj)
} }
if (&obj == DrawVisibleTrain) if (&obj == DrawVisibleTrain)
{ {
upDownPass(getTraining()); upDownPass(getTrainingPattern());
} }
if (&obj == DrawContextTrain) if (&obj == DrawContextTrain)
{ {
upDownPass(getTraining()); upDownPass(getTrainingPattern());
} }
} }
@@ -249,14 +249,14 @@ void RbmComponent::buttonClicked(Button* buttonThatWasClicked)
else if (buttonThatWasClicked == m_buttonCopyH2C) else if (buttonThatWasClicked == m_buttonCopyH2C)
{ {
DrawContextTrain->getData() = DrawHidden->getData(); DrawContextTrain->getData() = DrawHidden->getData();
upDownPass(getTraining()); upDownPass(getTrainingPattern());
} }
} }
void RbmComponent::redrawReconstruction() void RbmComponent::redrawReconstruction()
{ {
RbmComponent *pComp = static_cast<RbmComponent*> (root()); RbmComponent *pComp = static_cast<RbmComponent*> (root());
pComp->upDownPass(pComp->getTraining()); pComp->upDownPass(pComp->getTrainingPattern());
} }
void RbmComponent::gibbs(const arma::mat& vc) void RbmComponent::gibbs(const arma::mat& vc)
@@ -277,7 +277,7 @@ void RbmComponent::gibbs(const arma::mat& vc)
DrawContextReconst->DrawData(); DrawContextReconst->DrawData();
} }
arma::mat RbmComponent::getTraining() const arma::mat RbmComponent::getTrainingPattern() const
{ {
return arma::join_rows(DrawVisibleTrain->getData(), DrawContextTrain->getData()); return arma::join_rows(DrawVisibleTrain->getData(), DrawContextTrain->getData());
} }
+1 -1
View File
@@ -73,7 +73,7 @@ public:
ScopedPointer<DrawComponent> DrawVisibleTrain; ScopedPointer<DrawComponent> DrawVisibleTrain;
ScopedPointer<DrawComponent> DrawHidden; ScopedPointer<DrawComponent> DrawHidden;
ScopedPointer<DrawComponent> DrawContextTrain; ScopedPointer<DrawComponent> DrawContextTrain;
arma::mat getTraining() const; arma::mat getTrainingPattern() const;
arma::mat getReconst() const; arma::mat getReconst() const;
void trainRedraw(const arma::mat& vc); void trainRedraw(const arma::mat& vc);
void reconstRedraw(const arma::mat& vc); void reconstRedraw(const arma::mat& vc);
+44 -22
View File
@@ -175,7 +175,7 @@ bool Stack::loadWeights()
Layer *pLayer = m_pLayers; Layer *pLayer = m_pLayers;
while(pLayer) while(pLayer)
{ {
if (!pLayer->loadWeights(m_dir + "/" + m_name)) if (!pLayer->weightsLoad(m_dir, m_name))
{ {
return false; return false;
} }
@@ -189,7 +189,7 @@ bool Stack::saveWeights()
Layer *pLayer = m_pLayers; Layer *pLayer = m_pLayers;
while(pLayer) while(pLayer)
{ {
if (!pLayer->saveWeights(m_dir + "/" + m_name)) if (!pLayer->weightsSave(m_dir, m_name))
{ {
return false; return false;
} }
@@ -204,20 +204,20 @@ void Stack::train(Rbm::IListener* pListener)
while(pLayer) while(pLayer)
{ {
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl; 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->train(pListener);
pLayer = pLayer->next; 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; Layer *pLayer = m_pLayers;
while (pLayer) while (pLayer)
{ {
@@ -231,8 +231,19 @@ arma::mat Stack::trainingData(Layer* pThatLayer)
return thisBatch; 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 numTraining = 0;
uint32_t numVisible = 0; uint32_t numVisible = 0;
@@ -255,7 +266,7 @@ size_t Stack::loadTraining(bool doNormalize)
{ {
return 0; return 0;
} }
m_trainingData = arma::zeros(numTraining, numVisible); m_trainingBatch = arma::zeros(numTraining, numVisible);
uint32_t i, j; uint32_t i, j;
for (i=0; i < numTraining; i++) for (i=0; i < numTraining; i++)
@@ -266,7 +277,7 @@ size_t Stack::loadTraining(bool doNormalize)
int result = fscanf(pFile, "%f", &v); int result = fscanf(pFile, "%f", &v);
if (result > 0) 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) if (doNormalize)
{ {
m_trainingData = Rbm::normalize(m_trainingData); m_trainingBatch = Rbm::normalize(m_trainingData);
} }
return numTraining; 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"; std::string filename = m_dir + "/" + m_name + ".training.dat";
FILE *pFile = fopen(filename.c_str(), "w"); FILE *pFile = fopen(filename.c_str(), "w");
@@ -291,37 +313,37 @@ size_t Stack::saveTraining()
return 0; return 0;
} }
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_rows); fprintf(pFile, "%u\n", (uint32_t)m_trainingBatch.n_rows);
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_cols); fprintf(pFile, "%u\n", (uint32_t)m_trainingBatch.n_cols);
uint32_t i, j; 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); fclose(pFile);
std::cout << "Saved " << m_trainingData.n_rows << " training samples\n"; std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
return m_trainingData.n_rows; return m_trainingBatch.n_rows;
} }
size_t Stack::numTraining() size_t Stack::numTraining()
{ {
return m_trainingData.n_rows; return m_trainingBatch.n_rows;
} }
void Stack::addTraining(const arma::mat &toAdd) 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) void Stack::delTraining(int index)
{ {
m_trainingData.shed_row(index); m_trainingBatch.shed_row(index);
} }
+5 -5
View File
@@ -55,16 +55,16 @@ public:
size_t numTraining(); size_t numTraining();
void addTraining(const arma::mat &toAdd); void addTraining(const arma::mat &toAdd);
void delTraining(int index); void delTraining(int index);
size_t loadTraining(bool doNormalize=false); size_t loadTrainingBatch(bool doNormalize=false);
size_t saveTraining(); size_t saveTrainingBatch();
arma::mat& trainingData(); arma::mat& trainingBatch();
arma::mat trainingData(Layer *pLayer); arma::mat trainingBatch(Layer *pLayer);
private: private:
std::string m_dir; std::string m_dir;
std::string m_name; std::string m_name;
Layer *m_pLayers; Layer *m_pLayers;
arma::mat m_trainingData; arma::mat m_trainingBatch;
}; };
+5 -5
View File
@@ -82,13 +82,13 @@ int main()
RbmListener statusDisplay; RbmListener statusDisplay;
Stack stack(project); Stack stack(project);
stack.loadTraining(); stack.loadTrainingBatch();
stack.addTraining(stack.trainingData().row(1)); stack.addTraining(stack.trainingBatch().row(1));
printf("There are %d training samples\n", (int)stack.trainingData().n_rows); printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
stack.delTraining(0); stack.delTraining(0);
printf("There are %d training samples\n", (int)stack.trainingData().n_rows); printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
#if 1 #if 1
const int numLayers = 4; const int numLayers = 4;
@@ -129,7 +129,7 @@ int main()
stack.saveWeights(); stack.saveWeights();
Layer *layer = stack.getLayer(0); Layer *layer = stack.getLayer(0);
arma::mat v = arma::randu(stack.trainingData().n_rows, layer->bv().n_elem); arma::mat v = arma::randu(stack.trainingBatch().n_rows, layer->bv().n_elem);
arma::mat h = layer->toHiddenProbs(v); arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h); arma::mat r = layer->toVisibleProbs(h);
return 0; return 0;