- added setBatch()

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@771 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-10 10:36:55 +00:00
parent e240eae073
commit 648d8a002e
6 changed files with 19 additions and 8 deletions
+1
View File
@@ -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
View File
@@ -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;
};
+2 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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;
}
}