- added Stack::loadTraining(), Stack::addTraing(), Stack::delTraining()
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@623 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
@@ -61,6 +61,11 @@ void Stack::addLayer(Layer *pOtherLayer)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
void Stack::delLayer(Layer* pLayer)
|
||||||
|
{
|
||||||
|
assert(!"Stack::delLayer: Not implemented!");
|
||||||
|
}
|
||||||
|
|
||||||
Layer* Stack::getLayer(size_t layerId) const
|
Layer* Stack::getLayer(size_t layerId) const
|
||||||
{
|
{
|
||||||
Layer *pLayer = m_pLayers;
|
Layer *pLayer = m_pLayers;
|
||||||
@@ -208,3 +213,62 @@ void Stack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListe
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
arma::mat Stack::loadTraining()
|
||||||
|
{
|
||||||
|
uint32_t numTraining = 0;
|
||||||
|
uint32_t numVisible = 0;
|
||||||
|
FILE *pFile;
|
||||||
|
|
||||||
|
std::string filename = m_name + ".training.dat";
|
||||||
|
pFile = fopen(filename.c_str(), "r");
|
||||||
|
|
||||||
|
if (!pFile)
|
||||||
|
{
|
||||||
|
std::cout << "Could not open " << filename << "!" << std::endl;
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
int result = fscanf(pFile, "%d\n", &numTraining);
|
||||||
|
if (result < 0)
|
||||||
|
{
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
result = fscanf(pFile, "%d\n", &numVisible);
|
||||||
|
if (result < 0)
|
||||||
|
{
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
arma::mat data = arma::zeros(numTraining, numVisible);
|
||||||
|
|
||||||
|
uint32_t i, j;
|
||||||
|
for (i=0; i < numTraining; i++)
|
||||||
|
{
|
||||||
|
for (j=0; j < numVisible; j++)
|
||||||
|
{
|
||||||
|
float v;
|
||||||
|
int result = fscanf(pFile, "%f", &v);
|
||||||
|
if (result > 0)
|
||||||
|
{
|
||||||
|
data(i, j) = v;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fclose(pFile);
|
||||||
|
return data;
|
||||||
|
}
|
||||||
|
|
||||||
|
size_t Stack::numTraining(const arma::mat &batch)
|
||||||
|
{
|
||||||
|
return batch.n_rows;
|
||||||
|
}
|
||||||
|
|
||||||
|
void Stack::addTraining(arma::mat &batch, const arma::mat &toAdd)
|
||||||
|
{
|
||||||
|
batch.insert_rows(batch.n_rows, toAdd);
|
||||||
|
}
|
||||||
|
|
||||||
|
void Stack::delTraining(arma::mat &batch, int index)
|
||||||
|
{
|
||||||
|
batch.shed_row(index);
|
||||||
|
}
|
||||||
|
|
||||||
|
|||||||
+6
-6
@@ -39,7 +39,9 @@ public:
|
|||||||
Stack(const Stack& orig);
|
Stack(const Stack& orig);
|
||||||
virtual ~Stack();
|
virtual ~Stack();
|
||||||
|
|
||||||
|
size_t numLayers();
|
||||||
void addLayer(Layer *pLayer);
|
void addLayer(Layer *pLayer);
|
||||||
|
void delLayer(Layer *pLayer);
|
||||||
Layer* getLayer(size_t layerId) const;
|
Layer* getLayer(size_t layerId) const;
|
||||||
|
|
||||||
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener);
|
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener);
|
||||||
@@ -50,12 +52,10 @@ public:
|
|||||||
bool loadWeights();
|
bool loadWeights();
|
||||||
bool saveWeights();
|
bool saveWeights();
|
||||||
|
|
||||||
void batchchanged()
|
static size_t numTraining(const arma::mat &batch);
|
||||||
{
|
static void addTraining(arma::mat &batch, const arma::mat &toAdd);
|
||||||
std::cout << "Stack::batchchanged() called" << std::endl;
|
static void delTraining(arma::mat &batch, int index);
|
||||||
}
|
arma::mat loadTraining();
|
||||||
|
|
||||||
size_t numLayers();
|
|
||||||
|
|
||||||
private:
|
private:
|
||||||
std::string m_name;
|
std::string m_name;
|
||||||
|
|||||||
+9
-4
@@ -84,10 +84,15 @@ int main()
|
|||||||
RbmListener statusDisplay;
|
RbmListener statusDisplay;
|
||||||
Stack stack(project);
|
Stack stack(project);
|
||||||
|
|
||||||
arma::mat batch = loadTraining(project + string(".training.dat"));
|
arma::mat batch = stack.loadTraining();
|
||||||
|
|
||||||
size_t numTraining = batch.n_rows;
|
printf("Loaded %d training samples\n", (int)batch.n_rows);
|
||||||
printf("Loaded %d training samples\n", (int)numTraining);
|
|
||||||
|
stack.addTraining(batch, batch.row(1));
|
||||||
|
printf("Loaded %d training samples\n", (int)batch.n_rows);
|
||||||
|
|
||||||
|
stack.delTraining(batch, 0);
|
||||||
|
printf("Loaded %d training samples\n", (int)batch.n_rows);
|
||||||
|
|
||||||
#if 1
|
#if 1
|
||||||
const int numLayers = 4;
|
const int numLayers = 4;
|
||||||
@@ -128,7 +133,7 @@ int main()
|
|||||||
stack.saveWeights();
|
stack.saveWeights();
|
||||||
|
|
||||||
Layer *layer = stack.getLayer(0);
|
Layer *layer = stack.getLayer(0);
|
||||||
arma::mat v = arma::randu(numTraining, layer->bv().n_elem);
|
arma::mat v = arma::randu(batch.n_rows, layer->bv().n_elem);
|
||||||
arma::mat h = layer->toHiddenProbs(v);
|
arma::mat h = layer->toHiddenProbs(v);
|
||||||
arma::mat r = layer->toVisibleProbs(h);
|
arma::mat r = layer->toVisibleProbs(h);
|
||||||
return 0;
|
return 0;
|
||||||
|
|||||||
Reference in New Issue
Block a user