- 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_numVisibleY(numVisibleY)
|
||||
, m_numContext(numContext)
|
||||
, m_context(1, numContext)
|
||||
{
|
||||
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
|
||||
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
|
||||
|
||||
+5
-4
@@ -33,8 +33,8 @@ public:
|
||||
Json::Value toJson() const;
|
||||
bool loadWeights(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)
|
||||
{
|
||||
@@ -49,11 +49,11 @@ public:
|
||||
c_states.row(i) = ctx;
|
||||
}
|
||||
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
|
||||
{
|
||||
Rbm::train(trainingData(batch), pListener);
|
||||
Rbm::setBatch(trainingData(batch));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -136,6 +136,7 @@ private:
|
||||
size_t m_numVisibleX;
|
||||
size_t m_numVisibleY;
|
||||
size_t m_numContext;
|
||||
arma::mat m_context;
|
||||
|
||||
};
|
||||
|
||||
|
||||
@@ -899,7 +899,8 @@ const juce::String& MainComponent::getBaseDir()
|
||||
void MainComponent::run()
|
||||
{
|
||||
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"));
|
||||
}
|
||||
|
||||
|
||||
+2
-1
@@ -116,9 +116,10 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::
|
||||
dbhv -= sum(h_probs, 0);
|
||||
}
|
||||
|
||||
void Rbm::train(const arma::mat& batch, IListener* pListener)
|
||||
void Rbm::train(IListener* pListener)
|
||||
{
|
||||
Status status;
|
||||
arma::mat const &batch = m_batch;
|
||||
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
|
||||
double progress = 0;
|
||||
int lastProgress = -100;
|
||||
|
||||
+7
-1
@@ -126,7 +126,12 @@ public:
|
||||
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);
|
||||
const arma::mat& whv() const;
|
||||
@@ -179,6 +184,7 @@ protected:
|
||||
private:
|
||||
noise_gen_t m_noise;
|
||||
arma::mat m_whv;
|
||||
arma::mat m_batch;
|
||||
|
||||
};
|
||||
|
||||
|
||||
+2
-1
@@ -204,7 +204,8 @@ void Stack::train(Rbm::IListener* pListener)
|
||||
while(pLayer)
|
||||
{
|
||||
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;
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user