- 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;