- return name()

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@633 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-07 19:25:09 +00:00
parent 31ec578944
commit 732f6b67f6
2 changed files with 12 additions and 5 deletions
+7 -5
View File
@@ -52,7 +52,7 @@ bool Layer::loadWeights(const string &prjname)
string filename = m_weightsFile; string filename = m_weightsFile;
if (prjname.size() > 0) if (prjname.size() > 0)
{ {
filename = prjname + string(".") + m_weightsFile; filename = prjname + "." + m_weightsFile;
} }
FILE *pFile = fopen(filename.c_str(),"r"); FILE *pFile = fopen(filename.c_str(),"r");
@@ -62,6 +62,7 @@ bool Layer::loadWeights(const string &prjname)
return false; return false;
} }
std::cout << "Importing weights for " << m_name << "." << to_string((int)m_id) << std::endl;
int result = fscanf(pFile, "%d %d %d\n", &numVisibleX, &numVisibleY, &numHidden); int result = fscanf(pFile, "%d %d %d\n", &numVisibleX, &numVisibleY, &numHidden);
if (result < 0) if (result < 0)
{ {
@@ -107,10 +108,12 @@ bool Layer::loadWeights(const string &prjname)
bool Layer::saveWeights(const string &prjname) bool Layer::saveWeights(const string &prjname)
{ {
int numHidden = m_bh.n_elem;
int numVisible = m_bv.n_elem;
string filename = m_weightsFile; string filename = m_weightsFile;
if (prjname.size() > 0) if (prjname.size() > 0)
{ {
filename = prjname + string(".") + m_weightsFile; filename = prjname + "." + m_weightsFile;
} }
FILE *pFile = fopen(filename.c_str(),"w"); FILE *pFile = fopen(filename.c_str(),"w");
@@ -120,10 +123,9 @@ bool Layer::saveWeights(const string &prjname)
return false; return false;
} }
int numHidden = m_bh.n_elem; std::cout << "Exporting weights for " << m_name << "." << to_string((int)m_id) << std::endl;
int numVisible = m_bv.n_elem;
fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, numHidden);
fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, numHidden);
int i, j; int i, j;
for (i=0; i < numVisible; i++) for (i=0; i < numVisible; i++)
+5
View File
@@ -37,6 +37,11 @@ public:
arma::mat up_pass(const arma::mat& hidden); arma::mat up_pass(const arma::mat& hidden);
arma::mat down_pass(const arma::mat& visible); arma::mat down_pass(const arma::mat& visible);
std::string& name()
{
return m_name;
}
size_t id() size_t id()
{ {
return m_id; return m_id;