- 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
+6 -3
View File
@@ -132,12 +132,15 @@ public:
private:
const Params &m_params;
arma::mat m_w;
arma::mat m_bh;
arma::mat m_bv;
arma::mat sample(arma::mat const &src);
static arma::mat probsLogistic(arma::mat const &src);
void uniform(arma::mat &srcDst, double mu=0.0, double stdDev=1.0);
protected:
arma::mat m_w;
arma::mat m_bh;
arma::mat m_bv;
};
#endif /* RBM_HPP */
+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
+2 -1
View File
@@ -32,7 +32,8 @@ public:
virtual ~RbmLayer();
Json::Value toJson() const;
void saveWeights();
bool loadWeights();
bool saveWeights();
arma::mat up_pass(const arma::mat& hidden);
arma::mat down_pass(const arma::mat& visible);
+23 -3
View File
@@ -65,7 +65,7 @@ RbmLayer* Stack::getLayer(size_t layerId) const
}
void Stack::save(size_t numTraining)
bool Stack::save(size_t numTraining)
{
std::cout << "Exporting Project " << m_prjname << std::endl;
ofstream ofs(m_prjname + string(".prj"));
@@ -86,16 +86,36 @@ void Stack::save(size_t numTraining)
project["stack"]["layers"] = layers;
ofs << writer.write(project);
return true;
}
void Stack::saveWeights()
bool Stack::loadWeights()
{
RbmLayer *pLayer = m_pLayers;
while(pLayer)
{
pLayer->saveWeights();
if (!pLayer->loadWeights())
{
return false;
}
pLayer = pLayer->upper;
}
return true;
}
bool Stack::saveWeights()
{
RbmLayer *pLayer = m_pLayers;
while(pLayer)
{
if (!pLayer->saveWeights())
{
return false;
}
pLayer = pLayer->upper;
}
return true;
}
void Stack::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener)
+3 -2
View File
@@ -32,8 +32,9 @@ public:
void train(size_t layerId, const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener);
void train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener);
void save(size_t numTraining);
void saveWeights();
bool save(size_t numTraining);
bool loadWeights();
bool saveWeights();
private:
const std::string &m_prjname;
+9
View File
@@ -107,8 +107,17 @@ int main()
stack.addLayer(layer);
numHidden >>= 1;
}
// Save project
stack.save(numTraining);
// Load weights
stack.loadWeights();
// Train stack
stack.train(batch, 1000, 100, &statusDisplay);
// Save weights
stack.saveWeights();
RbmLayer *layer = stack.getLayer(0);