- added setBatch()
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@771 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
@@ -23,6 +23,7 @@ Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibl
|
|||||||
, m_numVisibleX(numVisibleX)
|
, m_numVisibleX(numVisibleX)
|
||||||
, m_numVisibleY(numVisibleY)
|
, m_numVisibleY(numVisibleY)
|
||||||
, m_numContext(numContext)
|
, m_numContext(numContext)
|
||||||
|
, m_context(1, numContext)
|
||||||
{
|
{
|
||||||
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
|
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
|
||||||
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
|
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
|
||||||
|
|||||||
+5
-4
@@ -33,8 +33,8 @@ public:
|
|||||||
Json::Value toJson() const;
|
Json::Value toJson() const;
|
||||||
bool loadWeights(const std::string &prjname="");
|
bool loadWeights(const std::string &prjname="");
|
||||||
bool saveWeights(const std::string &prjname="");
|
bool saveWeights(const std::string &prjname="");
|
||||||
|
|
||||||
void train(arma::mat const &batch, IListener *pListener=nullptr)
|
void setBatch(arma::mat const &batch)
|
||||||
{
|
{
|
||||||
if (batch.n_rows > 0)
|
if (batch.n_rows > 0)
|
||||||
{
|
{
|
||||||
@@ -49,11 +49,11 @@ public:
|
|||||||
c_states.row(i) = ctx;
|
c_states.row(i) = ctx;
|
||||||
}
|
}
|
||||||
arma::mat batch_with_ctx = arma::join_rows(batch, c_states);
|
arma::mat batch_with_ctx = arma::join_rows(batch, c_states);
|
||||||
Rbm::train(trainingData(batch_with_ctx), pListener);
|
Rbm::setBatch(trainingData(batch_with_ctx));
|
||||||
}
|
}
|
||||||
else
|
else
|
||||||
{
|
{
|
||||||
Rbm::train(trainingData(batch), pListener);
|
Rbm::setBatch(trainingData(batch));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -136,6 +136,7 @@ private:
|
|||||||
size_t m_numVisibleX;
|
size_t m_numVisibleX;
|
||||||
size_t m_numVisibleY;
|
size_t m_numVisibleY;
|
||||||
size_t m_numContext;
|
size_t m_numContext;
|
||||||
|
arma::mat m_context;
|
||||||
|
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -899,7 +899,8 @@ const juce::String& MainComponent::getBaseDir()
|
|||||||
void MainComponent::run()
|
void MainComponent::run()
|
||||||
{
|
{
|
||||||
trainButton->setButtonText (TRANS("Stop"));
|
trainButton->setButtonText (TRANS("Stop"));
|
||||||
m_pLayer->train(m_stack->trainingData(), this);
|
m_pLayer->setBatch(m_stack->trainingData());
|
||||||
|
m_pLayer->train(this);
|
||||||
trainButton->setButtonText (TRANS("Train"));
|
trainButton->setButtonText (TRANS("Train"));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+2
-1
@@ -116,9 +116,10 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::
|
|||||||
dbhv -= sum(h_probs, 0);
|
dbhv -= sum(h_probs, 0);
|
||||||
}
|
}
|
||||||
|
|
||||||
void Rbm::train(const arma::mat& batch, IListener* pListener)
|
void Rbm::train(IListener* pListener)
|
||||||
{
|
{
|
||||||
Status status;
|
Status status;
|
||||||
|
arma::mat const &batch = m_batch;
|
||||||
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
|
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
|
||||||
double progress = 0;
|
double progress = 0;
|
||||||
int lastProgress = -100;
|
int lastProgress = -100;
|
||||||
|
|||||||
+7
-1
@@ -126,7 +126,12 @@ public:
|
|||||||
m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv;
|
m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv;
|
||||||
}
|
}
|
||||||
|
|
||||||
void train(arma::mat const &batch, IListener *pListener=nullptr);
|
void setBatch(arma::mat const &batch)
|
||||||
|
{
|
||||||
|
m_batch = batch;
|
||||||
|
}
|
||||||
|
|
||||||
|
void train(IListener *pListener=nullptr);
|
||||||
|
|
||||||
static arma::mat normalize(const arma::mat &hidden);
|
static arma::mat normalize(const arma::mat &hidden);
|
||||||
const arma::mat& whv() const;
|
const arma::mat& whv() const;
|
||||||
@@ -179,6 +184,7 @@ protected:
|
|||||||
private:
|
private:
|
||||||
noise_gen_t m_noise;
|
noise_gen_t m_noise;
|
||||||
arma::mat m_whv;
|
arma::mat m_whv;
|
||||||
|
arma::mat m_batch;
|
||||||
|
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
+2
-1
@@ -204,7 +204,8 @@ void Stack::train(Rbm::IListener* pListener)
|
|||||||
while(pLayer)
|
while(pLayer)
|
||||||
{
|
{
|
||||||
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
|
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
|
||||||
pLayer->train(m_trainingData, pListener);
|
pLayer->setBatch(m_trainingData);
|
||||||
|
pLayer->train(pListener);
|
||||||
pLayer = pLayer->next;
|
pLayer = pLayer->next;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user