- 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:
|
private:
|
||||||
const Params &m_params;
|
const Params &m_params;
|
||||||
arma::mat m_w;
|
|
||||||
arma::mat m_bh;
|
|
||||||
arma::mat m_bv;
|
|
||||||
arma::mat sample(arma::mat const &src);
|
arma::mat sample(arma::mat const &src);
|
||||||
static arma::mat probsLogistic(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);
|
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 */
|
#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)
|
if (!pFile)
|
||||||
{
|
{
|
||||||
std::cout << "Could not open " << m_weightsFile << "!" << std::endl;
|
std::cout << "loadWeights(): Could not open " << m_weightsFile << " for reading!" << std::endl;
|
||||||
return;
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
size_t numHidden = bh().n_elem;
|
int result = fscanf(pFile, "%d %d %d\n", &numVisibleX, &numVisibleY, &numHidden);
|
||||||
size_t numVisible = bh().n_elem;
|
if (result < 0)
|
||||||
fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, (int)numHidden);
|
{
|
||||||
|
return false;
|
||||||
uint32_t i, j;
|
}
|
||||||
|
|
||||||
|
size_t numVisible = numVisibleX*numVisibleY;
|
||||||
|
|
||||||
|
int i, j;
|
||||||
|
float v;
|
||||||
for (i=0; i < numVisible; i++)
|
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++)
|
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 (i=0; i < numVisible; i++)
|
||||||
{
|
{
|
||||||
for (j=0; j < numHidden; j++)
|
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");
|
fprintf(pFile, "\n");
|
||||||
}
|
}
|
||||||
|
|
||||||
fclose(pFile);
|
fclose(pFile);
|
||||||
|
|
||||||
|
return true;
|
||||||
}
|
}
|
||||||
|
|
||||||
Json::Value RbmLayer::toJson() const
|
Json::Value RbmLayer::toJson() const
|
||||||
|
|||||||
+2
-1
@@ -32,7 +32,8 @@ public:
|
|||||||
virtual ~RbmLayer();
|
virtual ~RbmLayer();
|
||||||
|
|
||||||
Json::Value toJson() const;
|
Json::Value toJson() const;
|
||||||
void saveWeights();
|
bool loadWeights();
|
||||||
|
bool saveWeights();
|
||||||
|
|
||||||
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);
|
||||||
|
|||||||
+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;
|
std::cout << "Exporting Project " << m_prjname << std::endl;
|
||||||
ofstream ofs(m_prjname + string(".prj"));
|
ofstream ofs(m_prjname + string(".prj"));
|
||||||
@@ -86,16 +86,36 @@ void Stack::save(size_t numTraining)
|
|||||||
project["stack"]["layers"] = layers;
|
project["stack"]["layers"] = layers;
|
||||||
|
|
||||||
ofs << writer.write(project);
|
ofs << writer.write(project);
|
||||||
|
|
||||||
|
return true;
|
||||||
}
|
}
|
||||||
|
|
||||||
void Stack::saveWeights()
|
bool Stack::loadWeights()
|
||||||
{
|
{
|
||||||
RbmLayer *pLayer = m_pLayers;
|
RbmLayer *pLayer = m_pLayers;
|
||||||
while(pLayer)
|
while(pLayer)
|
||||||
{
|
{
|
||||||
pLayer->saveWeights();
|
if (!pLayer->loadWeights())
|
||||||
|
{
|
||||||
|
return false;
|
||||||
|
}
|
||||||
pLayer = pLayer->upper;
|
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)
|
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(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 train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener);
|
||||||
void save(size_t numTraining);
|
bool save(size_t numTraining);
|
||||||
void saveWeights();
|
bool loadWeights();
|
||||||
|
bool saveWeights();
|
||||||
|
|
||||||
private:
|
private:
|
||||||
const std::string &m_prjname;
|
const std::string &m_prjname;
|
||||||
|
|||||||
@@ -107,8 +107,17 @@ int main()
|
|||||||
stack.addLayer(layer);
|
stack.addLayer(layer);
|
||||||
numHidden >>= 1;
|
numHidden >>= 1;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Save project
|
||||||
stack.save(numTraining);
|
stack.save(numTraining);
|
||||||
|
|
||||||
|
// Load weights
|
||||||
|
stack.loadWeights();
|
||||||
|
|
||||||
|
// Train stack
|
||||||
stack.train(batch, 1000, 100, &statusDisplay);
|
stack.train(batch, 1000, 100, &statusDisplay);
|
||||||
|
|
||||||
|
// Save weights
|
||||||
stack.saveWeights();
|
stack.saveWeights();
|
||||||
|
|
||||||
RbmLayer *layer = stack.getLayer(0);
|
RbmLayer *layer = stack.getLayer(0);
|
||||||
|
|||||||
Reference in New Issue
Block a user