- 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)
{
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)
@@ -34,7 +33,6 @@ Layer::Layer(const Layer& orig)
, next(nullptr)
, prev(nullptr)
, m_name(orig.m_name)
, m_weightsFile(orig.m_weightsFile)
, m_id(orig.m_id)
, m_numVisibleX(orig.m_numVisibleX)
, m_numVisibleY(orig.m_numVisibleY)
@@ -45,13 +43,10 @@ Layer::~Layer()
{
}
bool Layer::loadWeights(const string &prjname)
{
string filename = m_weightsFile;
if (prjname.size() > 0)
{
filename = prjname + "." + m_weightsFile;
}
string filename = filePrefix(prjname) + ".weights.dat";
FILE *pFile = fopen(filename.c_str(),"r");
if (!pFile)
@@ -108,11 +103,7 @@ bool Layer::saveWeights(const string &prjname)
{
int numHidden = m_bhv.n_elem;
int numVisible = m_bv.n_elem;
string filename = m_weightsFile;
if (prjname.size() > 0)
{
filename = prjname + "." + m_weightsFile;
}
string filename = filePrefix(prjname) + ".weights.dat";
FILE *pFile = fopen(filename.c_str(),"w");
if (!pFile)
@@ -153,7 +144,7 @@ Json::Value Layer::toJson() const
Json::Value layer;
layer["name"] = m_name;
layer["id"] = (int)m_id;
layer["weights_file"] = m_weightsFile;
layer["weights_file"] = filePrefix("") + ".weights.dat";
layer["numVisibleX"] = (int)m_numVisibleX;
layer["numVisibleY"] = (int)m_numVisibleY;
layer["numHidden"] = (int)whv().n_cols;
+51 -4
View File
@@ -31,9 +31,7 @@ public:
virtual ~Layer();
Json::Value toJson() const;
bool loadWeights(const std::string &prjname="");
bool saveWeights(const std::string &prjname="");
void setBatch(arma::mat const &batch)
{
if (batch.n_rows > 0)
@@ -83,6 +81,43 @@ public:
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 thisBatch = batch;
@@ -131,12 +166,24 @@ public:
private:
std::string m_name;
std::string m_weightsFile;
size_t m_id;
size_t m_numVisibleX;
size_t m_numVisibleY;
size_t m_numContext;
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;
}
};
+4 -4
View File
@@ -814,7 +814,7 @@ void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged)
m_pLayer = static_cast<RbmComponent*>(m_stack->getLayer(index));
updateControls();
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));
}
@@ -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...
void MainComponent::save ()
{
m_stack->saveTraining();
m_stack->saveTrainingBatch();
m_stack->saveWeights();
m_stack->save();
}
@@ -899,7 +899,7 @@ const juce::String& MainComponent::getBaseDir()
void MainComponent::run()
{
trainButton->setButtonText (TRANS("Stop"));
m_pLayer->setBatch(m_stack->trainingData());
m_pLayer->setBatch(m_stack->trainingBatch());
m_pLayer->train(this);
trainButton->setButtonText (TRANS("Train"));
}
@@ -916,7 +916,7 @@ bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
}
RbmComponent *pComp = static_cast<RbmComponent*>(pRbm);
pComp->upPass(pComp->getTraining());
pComp->upPass(pComp->getTrainingPattern());
pComp->redrawReconstruction();
pComp->redrawWeights();
+5 -5
View File
@@ -98,12 +98,12 @@ private:
void clearTraining()
{
patterSlider->setRange(0, 0, 1);
m_stack->trainingData().clear();
m_stack->trainingBatch().clear();
}
void loadTraining()
{
m_stack->loadTraining(rbmNormalizeDataToggleButton->getToggleState());
m_stack->loadTrainingBatch(rbmNormalizeDataToggleButton->getToggleState());
patterSlider->setRange(0, m_stack->numTraining()-1, 1);
}
@@ -128,16 +128,16 @@ private:
{
if (m_pLayer->context().is_empty())
{
m_pLayer->setBatch(m_stack->trainingData());
m_pLayer->setBatch(m_stack->trainingBatch());
}
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
{
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)
{
upDownPass(getTraining());
upDownPass(getTrainingPattern());
}
if (&obj == DrawContextTrain)
{
upDownPass(getTraining());
upDownPass(getTrainingPattern());
}
}
@@ -249,14 +249,14 @@ void RbmComponent::buttonClicked(Button* buttonThatWasClicked)
else if (buttonThatWasClicked == m_buttonCopyH2C)
{
DrawContextTrain->getData() = DrawHidden->getData();
upDownPass(getTraining());
upDownPass(getTrainingPattern());
}
}
void RbmComponent::redrawReconstruction()
{
RbmComponent *pComp = static_cast<RbmComponent*> (root());
pComp->upDownPass(pComp->getTraining());
pComp->upDownPass(pComp->getTrainingPattern());
}
void RbmComponent::gibbs(const arma::mat& vc)
@@ -277,7 +277,7 @@ void RbmComponent::gibbs(const arma::mat& vc)
DrawContextReconst->DrawData();
}
arma::mat RbmComponent::getTraining() const
arma::mat RbmComponent::getTrainingPattern() const
{
return arma::join_rows(DrawVisibleTrain->getData(), DrawContextTrain->getData());
}
+1 -1
View File
@@ -73,7 +73,7 @@ public:
ScopedPointer<DrawComponent> DrawVisibleTrain;
ScopedPointer<DrawComponent> DrawHidden;
ScopedPointer<DrawComponent> DrawContextTrain;
arma::mat getTraining() const;
arma::mat getTrainingPattern() const;
arma::mat getReconst() const;
void trainRedraw(const arma::mat& vc);
void reconstRedraw(const arma::mat& vc);
+45 -23
View File
@@ -175,7 +175,7 @@ bool Stack::loadWeights()
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (!pLayer->loadWeights(m_dir + "/" + m_name))
if (!pLayer->weightsLoad(m_dir, m_name))
{
return false;
}
@@ -189,7 +189,7 @@ bool Stack::saveWeights()
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (!pLayer->saveWeights(m_dir + "/" + m_name))
if (!pLayer->weightsSave(m_dir, m_name))
{
return false;
}
@@ -204,20 +204,20 @@ void Stack::train(Rbm::IListener* pListener)
while(pLayer)
{
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 = 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;
while (pLayer)
{
@@ -231,8 +231,19 @@ arma::mat Stack::trainingData(Layer* pThatLayer)
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 numVisible = 0;
@@ -255,8 +266,8 @@ size_t Stack::loadTraining(bool doNormalize)
{
return 0;
}
m_trainingData = arma::zeros(numTraining, numVisible);
m_trainingBatch = arma::zeros(numTraining, numVisible);
uint32_t i, j;
for (i=0; i < numTraining; i++)
{
@@ -266,7 +277,7 @@ size_t Stack::loadTraining(bool doNormalize)
int result = fscanf(pFile, "%f", &v);
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)
{
m_trainingData = Rbm::normalize(m_trainingData);
m_trainingBatch = Rbm::normalize(m_trainingData);
}
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";
FILE *pFile = fopen(filename.c_str(), "w");
@@ -291,37 +313,37 @@ size_t Stack::saveTraining()
return 0;
}
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_rows);
fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_cols);
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_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);
std::cout << "Saved " << m_trainingData.n_rows << " training samples\n";
return m_trainingData.n_rows;
std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
return m_trainingBatch.n_rows;
}
size_t Stack::numTraining()
{
return m_trainingData.n_rows;
return m_trainingBatch.n_rows;
}
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)
{
m_trainingData.shed_row(index);
m_trainingBatch.shed_row(index);
}
+5 -5
View File
@@ -55,16 +55,16 @@ public:
size_t numTraining();
void addTraining(const arma::mat &toAdd);
void delTraining(int index);
size_t loadTraining(bool doNormalize=false);
size_t saveTraining();
arma::mat& trainingData();
arma::mat trainingData(Layer *pLayer);
size_t loadTrainingBatch(bool doNormalize=false);
size_t saveTrainingBatch();
arma::mat& trainingBatch();
arma::mat trainingBatch(Layer *pLayer);
private:
std::string m_dir;
std::string m_name;
Layer *m_pLayers;
arma::mat m_trainingData;
arma::mat m_trainingBatch;
};
+5 -5
View File
@@ -82,13 +82,13 @@ int main()
RbmListener statusDisplay;
Stack stack(project);
stack.loadTraining();
stack.loadTrainingBatch();
stack.addTraining(stack.trainingData().row(1));
printf("There are %d training samples\n", (int)stack.trainingData().n_rows);
stack.addTraining(stack.trainingBatch().row(1));
printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
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
const int numLayers = 4;
@@ -129,7 +129,7 @@ int main()
stack.saveWeights();
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 r = layer->toVisibleProbs(h);
return 0;