- Layer: removed legacy weights load/save
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@777 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+21
-82
@@ -43,99 +43,38 @@ Layer::~Layer()
|
|||||||
{
|
{
|
||||||
}
|
}
|
||||||
|
|
||||||
|
bool Layer::weightsLoad(std::string const &dir, std::string const &prj)
|
||||||
bool Layer::loadWeights(const string &prjname)
|
|
||||||
{
|
{
|
||||||
string filename = filePrefix(prjname) + ".weights.dat";
|
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);
|
||||||
|
|
||||||
FILE *pFile = fopen(filename.c_str(),"r");
|
if (result)
|
||||||
if (!pFile)
|
|
||||||
{
|
{
|
||||||
std::cout << "loadWeights(): Could not open " << filename << " for reading!" << std::endl;
|
std::cout << "Layer " << m_id << ": Importing weights" << std::endl;
|
||||||
return false;
|
weightsAssign(w, bh, bv);
|
||||||
}
|
}
|
||||||
|
|
||||||
std::cout << "Importing weights for " << m_name << "." << to_string((int)m_id) << std::endl;
|
return result;
|
||||||
|
|
||||||
size_t numVisible = Rbm::numVisible();
|
|
||||||
size_t numHidden = Rbm::numHidden();
|
|
||||||
int i, j;
|
|
||||||
float v;
|
|
||||||
int result;
|
|
||||||
arma::mat _bv(1, numVisible);
|
|
||||||
for (i=0; i < numVisible; i++)
|
|
||||||
{
|
|
||||||
result = fscanf(pFile, "%f", &v);
|
|
||||||
if (result > 0)
|
|
||||||
{
|
|
||||||
_bv(i) = v;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
arma::mat _bhv(1, numHidden);
|
|
||||||
for (i=0; i < numHidden; i++)
|
|
||||||
{
|
|
||||||
result = fscanf(pFile, "%f", &v);
|
|
||||||
if (result > 0)
|
|
||||||
{
|
|
||||||
_bhv(i) = v;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
arma::mat _whv(numVisible, numHidden);
|
|
||||||
for (i=0; i < numVisible; i++)
|
|
||||||
{
|
|
||||||
for (j=0; j < numHidden; j++)
|
|
||||||
{
|
|
||||||
|
|
||||||
result = fscanf(pFile, "%f", &v);
|
|
||||||
if (result > 0)
|
|
||||||
{
|
|
||||||
_whv(i, j) = v;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
weightsAssign(_whv, _bhv, _bv);
|
|
||||||
fclose(pFile);
|
|
||||||
|
|
||||||
return true;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
bool Layer::saveWeights(const string &prjname)
|
bool Layer::weightsSave(std::string const &dir, std::string const &prj)
|
||||||
{
|
{
|
||||||
int numHidden = m_bhv.n_elem;
|
bool result = true;
|
||||||
int numVisible = m_bv.n_elem;
|
result &= whv().save(filePrefix(prj) + ".w.dat", arma::arma_ascii);
|
||||||
string filename = filePrefix(prjname) + ".weights.dat";
|
result &= bh().save(filePrefix(prj) + ".bh.dat", arma::arma_ascii);
|
||||||
|
result &= bv().save(filePrefix(prj) + ".bv.dat", arma::arma_ascii);
|
||||||
|
|
||||||
FILE *pFile = fopen(filename.c_str(),"w");
|
if (result)
|
||||||
if (!pFile)
|
|
||||||
{
|
{
|
||||||
std::cout << "saveWeights(): Could not open " << filename << " for writing!" << std::endl;
|
std::cout << "Layer " << m_id << ": Exporting weights" << std::endl;
|
||||||
return false;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
std::cout << "Exporting weights for " << m_name << "." << to_string((int)m_id) << std::endl;
|
return result;
|
||||||
|
|
||||||
int i, j;
|
|
||||||
|
|
||||||
for (i=0; i < numVisible; i++)
|
|
||||||
{
|
|
||||||
fprintf(pFile, "%3.6f\n", m_bv(i));
|
|
||||||
}
|
|
||||||
for (i=0; i < numHidden; i++)
|
|
||||||
{
|
|
||||||
fprintf(pFile, "%3.6f\n", m_bhv(i));
|
|
||||||
}
|
|
||||||
const arma::mat &_whv = whv();
|
|
||||||
for (i=0; i < numVisible; i++)
|
|
||||||
{
|
|
||||||
for (j=0; j < numHidden; j++)
|
|
||||||
{
|
|
||||||
fprintf(pFile, "%3.6f ", _whv(i,j));
|
|
||||||
}
|
|
||||||
fprintf(pFile, "\n");
|
|
||||||
}
|
|
||||||
fclose(pFile);
|
|
||||||
|
|
||||||
return true;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Json::Value Layer::toJson() const
|
Json::Value Layer::toJson() const
|
||||||
|
|||||||
+2
-36
@@ -81,42 +81,8 @@ public:
|
|||||||
return whv().n_cols;
|
return whv().n_cols;
|
||||||
}
|
}
|
||||||
|
|
||||||
bool weightsLoad(std::string const &dir, std::string const &prj)
|
bool weightsLoad(std::string const &dir, std::string const &prj);
|
||||||
{
|
bool weightsSave(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)
|
||||||
{
|
{
|
||||||
|
|||||||
+13
-18
@@ -233,21 +233,19 @@ arma::mat Stack::trainingBatch(Layer* pThatLayer)
|
|||||||
|
|
||||||
size_t Stack::loadTrainingBatch(bool doNormalize)
|
size_t Stack::loadTrainingBatch(bool doNormalize)
|
||||||
{
|
{
|
||||||
{
|
std::string filename = m_dir + "/" + m_name + ".training.dat";
|
||||||
std::string path = m_dir + "/" + m_name + ".training.mat";
|
std::string path = m_dir + "/" + m_name + ".training.dat";
|
||||||
bool success = m_trainingBatch.load(path, arma::arma_ascii);
|
bool success = m_trainingBatch.load(filename, arma::arma_ascii);
|
||||||
|
|
||||||
if (success)
|
if (success)
|
||||||
{
|
{
|
||||||
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n";
|
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n";
|
||||||
return m_trainingBatch.n_rows;
|
return m_trainingBatch.n_rows;
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
uint32_t numTraining = 0;
|
uint32_t numTraining = 0;
|
||||||
uint32_t numVisible = 0;
|
uint32_t numVisible = 0;
|
||||||
|
|
||||||
std::string filename = m_dir + "/" + m_name + ".training.dat";
|
|
||||||
FILE *pFile = fopen(filename.c_str(), "r");
|
FILE *pFile = fopen(filename.c_str(), "r");
|
||||||
|
|
||||||
if (!pFile)
|
if (!pFile)
|
||||||
@@ -293,18 +291,15 @@ size_t Stack::loadTrainingBatch(bool doNormalize)
|
|||||||
|
|
||||||
size_t Stack::saveTrainingBatch()
|
size_t Stack::saveTrainingBatch()
|
||||||
{
|
{
|
||||||
{
|
std::string filename = m_dir + "/" + m_name + ".training.dat";
|
||||||
std::string path = m_dir + "/" + m_name + ".training.mat";
|
bool success = m_trainingBatch.save(filename, arma::arma_ascii);
|
||||||
bool success = m_trainingBatch.save(path, arma::arma_ascii);
|
|
||||||
|
|
||||||
if (success)
|
if (success)
|
||||||
{
|
{
|
||||||
std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
|
std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n";
|
||||||
return m_trainingBatch.n_rows;
|
return m_trainingBatch.n_rows;
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
std::string filename = m_dir + "/" + m_name + ".training.dat";
|
|
||||||
FILE *pFile = fopen(filename.c_str(), "w");
|
FILE *pFile = fopen(filename.c_str(), "w");
|
||||||
|
|
||||||
if (!pFile)
|
if (!pFile)
|
||||||
|
|||||||
Reference in New Issue
Block a user