- 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:
+4
-13
@@ -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
@@ -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;
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user