- cleaned up
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@821 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+15
-43
@@ -29,37 +29,37 @@ DeepStack::~DeepStack()
|
|||||||
{
|
{
|
||||||
}
|
}
|
||||||
|
|
||||||
void DeepStack::train(Rbm::IListener* pListener)
|
void DeepStack::train(const arma::mat& batch, Rbm::IListener* pListener)
|
||||||
{
|
{
|
||||||
Layer *pLayer = m_pLayers;
|
arma::mat thisBatch = batch;
|
||||||
|
Layer *pLayer = getLayer(0);
|
||||||
while (pLayer)
|
while (pLayer)
|
||||||
{
|
{
|
||||||
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
|
thisBatch = trainingBatchFrom(pLayer->id(), thisBatch);
|
||||||
pLayer->calcContextBatch(m_trainingBatch);
|
pLayer->train(thisBatch, pListener);
|
||||||
pLayer->train(m_trainingBatch, pListener);
|
|
||||||
pLayer = pLayer->next;
|
pLayer = pLayer->next;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
arma::mat& DeepStack::trainingBatch()
|
void DeepStack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener)
|
||||||
{
|
{
|
||||||
return m_trainingBatch;
|
arma::mat thisBatch = batch;
|
||||||
}
|
Layer *pLayer = getLayer(0);
|
||||||
|
|
||||||
arma::mat DeepStack::trainingBatch(Layer* pThatLayer)
|
|
||||||
{
|
|
||||||
arma::mat thisBatch = m_trainingBatch;
|
|
||||||
Layer *pLayer = m_pLayers;
|
|
||||||
while (pLayer)
|
while (pLayer)
|
||||||
{
|
{
|
||||||
if (pLayer->id() == pThatLayer->id())
|
if (pLayer->id() == layerId)
|
||||||
{
|
{
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
thisBatch = pLayer->toHiddenProbs(thisBatch);
|
thisBatch = pLayer->toHiddenProbs(thisBatch);
|
||||||
pLayer = pLayer->next;
|
pLayer = pLayer->next;
|
||||||
}
|
}
|
||||||
return thisBatch;
|
pLayer->train(thisBatch, pListener);
|
||||||
|
}
|
||||||
|
|
||||||
|
arma::mat& DeepStack::trainingBatch()
|
||||||
|
{
|
||||||
|
return m_trainingBatch;
|
||||||
}
|
}
|
||||||
|
|
||||||
size_t DeepStack::loadTrainingBatch(bool doNormalize)
|
size_t DeepStack::loadTrainingBatch(bool doNormalize)
|
||||||
@@ -123,34 +123,6 @@ void DeepStack::delTraining(int index)
|
|||||||
m_trainingBatch.shed_row(index);
|
m_trainingBatch.shed_row(index);
|
||||||
}
|
}
|
||||||
|
|
||||||
void DeepStack::train(const arma::mat& batch, Rbm::IListener* pListener)
|
|
||||||
{
|
|
||||||
arma::mat thisBatch = batch;
|
|
||||||
Layer *pLayer = getLayer(0);
|
|
||||||
while (pLayer)
|
|
||||||
{
|
|
||||||
pLayer->train(thisBatch, pListener);
|
|
||||||
thisBatch = pLayer->toHiddenProbs(thisBatch);
|
|
||||||
pLayer = pLayer->next;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
void DeepStack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener)
|
|
||||||
{
|
|
||||||
arma::mat thisBatch = batch;
|
|
||||||
Layer *pLayer = getLayer(0);
|
|
||||||
while (pLayer)
|
|
||||||
{
|
|
||||||
if (pLayer->id() == layerId)
|
|
||||||
{
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
thisBatch = pLayer->toHiddenProbs(thisBatch);
|
|
||||||
pLayer = pLayer->next;
|
|
||||||
}
|
|
||||||
pLayer->train(thisBatch, pListener);
|
|
||||||
}
|
|
||||||
|
|
||||||
arma::mat DeepStack::upPass(size_t layerId, const arma::mat& v)
|
arma::mat DeepStack::upPass(size_t layerId, const arma::mat& v)
|
||||||
{
|
{
|
||||||
arma::mat h = arma::zeros(0,0);
|
arma::mat h = arma::zeros(0,0);
|
||||||
|
|||||||
@@ -28,7 +28,6 @@ public:
|
|||||||
DeepStack(const DeepStack& orig);
|
DeepStack(const DeepStack& orig);
|
||||||
virtual ~DeepStack();
|
virtual ~DeepStack();
|
||||||
|
|
||||||
void train(Rbm::IListener* pListener);
|
|
||||||
void train(const arma::mat& batch, Rbm::IListener* pListener) override;
|
void train(const arma::mat& batch, Rbm::IListener* pListener) override;
|
||||||
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override;
|
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override;
|
||||||
|
|
||||||
@@ -42,7 +41,6 @@ public:
|
|||||||
size_t loadTrainingBatch(bool doNormalize=false);
|
size_t loadTrainingBatch(bool doNormalize=false);
|
||||||
size_t saveTrainingBatch();
|
size_t saveTrainingBatch();
|
||||||
arma::mat& trainingBatch();
|
arma::mat& trainingBatch();
|
||||||
arma::mat trainingBatch(Layer *pLayer);
|
|
||||||
|
|
||||||
private:
|
private:
|
||||||
arma::mat m_trainingBatch;
|
arma::mat m_trainingBatch;
|
||||||
|
|||||||
Reference in New Issue
Block a user