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