- 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:
+6
-3
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user