- 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 *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);
|
||||
virtual ~Stack();
|
||||
|
||||
size_t numLayers();
|
||||
void addLayer(Layer *pLayer);
|
||||
void delLayer(Layer *pLayer);
|
||||
Layer* getLayer(size_t layerId) const;
|
||||
|
||||
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener);
|
||||
@@ -50,12 +52,10 @@ public:
|
||||
bool loadWeights();
|
||||
bool saveWeights();
|
||||
|
||||
void batchchanged()
|
||||
{
|
||||
std::cout << "Stack::batchchanged() called" << std::endl;
|
||||
}
|
||||
|
||||
size_t numLayers();
|
||||
static size_t numTraining(const arma::mat &batch);
|
||||
static void addTraining(arma::mat &batch, const arma::mat &toAdd);
|
||||
static void delTraining(arma::mat &batch, int index);
|
||||
arma::mat loadTraining();
|
||||
|
||||
private:
|
||||
std::string m_name;
|
||||
|
||||
+9
-4
@@ -84,10 +84,15 @@ int main()
|
||||
RbmListener statusDisplay;
|
||||
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)numTraining);
|
||||
printf("Loaded %d training samples\n", (int)batch.n_rows);
|
||||
|
||||
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
|
||||
const int numLayers = 4;
|
||||
@@ -128,7 +133,7 @@ int main()
|
||||
stack.saveWeights();
|
||||
|
||||
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 r = layer->toVisibleProbs(h);
|
||||
return 0;
|
||||
|
||||
Reference in New Issue
Block a user