- 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_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
View File
@@ -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;
}; };
+2 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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;
} }
} }