- load and save of weights

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@588 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-28 17:26:59 +00:00
parent 27c07088c2
commit 9bfee1ae16
6 changed files with 114 additions and 22 deletions
+71 -13
View File
@@ -45,41 +45,99 @@ RbmLayer::~RbmLayer()
{
}
void RbmLayer::saveWeights()
bool RbmLayer::loadWeights()
{
FILE *pFile = fopen(m_weightsFile.c_str(), "w");
int numVisibleX;
int numVisibleY;
int numHidden;
FILE *pFile;
pFile = fopen(m_weightsFile.c_str(),"r");
if (!pFile)
{
std::cout << "Could not open " << m_weightsFile << "!" << std::endl;
return;
std::cout << "loadWeights(): Could not open " << m_weightsFile << " for reading!" << std::endl;
return false;
}
size_t numHidden = bh().n_elem;
size_t numVisible = bh().n_elem;
fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, (int)numHidden);
uint32_t i, j;
int result = fscanf(pFile, "%d %d %d\n", &numVisibleX, &numVisibleY, &numHidden);
if (result < 0)
{
return false;
}
size_t numVisible = numVisibleX*numVisibleY;
int i, j;
float v;
for (i=0; i < numVisible; i++)
{
fprintf(pFile, "%3.6f\n", bv()(i));
result = fscanf(pFile, "%f", &v);
if (result > 0)
{
m_bv(i) = v;
}
}
for (i=0; i < numHidden; i++)
{
fprintf(pFile, "%3.6f\n", bh()(i));
result = fscanf(pFile, "%f", &v);
if (result > 0)
{
m_bh(i) = v;
}
}
for (i=0; i < numVisible; i++)
{
for (j=0; j < numHidden; j++)
{
fprintf(pFile, "%3.6f ", w()(i,j));
result = fscanf(pFile, "%f", &v);
if (result > 0)
{
m_w(i, j) = v;
}
}
}
fclose(pFile);
return true;
}
bool RbmLayer::saveWeights()
{
FILE *pFile = fopen(m_weightsFile.c_str(), "w");
if (!pFile)
{
std::cout << "saveWeights(): Could not open " << m_weightsFile << " for writing!" << std::endl;
return false;
}
int numHidden = m_bh.n_elem;
int numVisible = m_bv.n_elem;
fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, numHidden);
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_bh(i));
}
for (i=0; i < numVisible; i++)
{
for (j=0; j < numHidden; j++)
{
fprintf(pFile, "%3.6f ", m_w(i,j));
}
fprintf(pFile, "\n");
}
fclose(pFile);
return true;
}
Json::Value RbmLayer::toJson() const