- return name()
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@633 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+7
-5
@@ -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++)
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
Reference in New Issue
Block a user