- integrated RnnStack

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@826 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-18 08:36:10 +00:00
parent 9b2861e3d8
commit 17770f402c
10 changed files with 130 additions and 34 deletions
+5 -1
View File
@@ -54,6 +54,10 @@ public:
StackType type();
void setName(const std::string &name);
size_t numLayers();
virtual size_t numContext()
{
return 0;
}
void addLayer(Layer *pLayer);
void delLayer(Layer *pLayer);
@@ -67,7 +71,7 @@ public:
virtual void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) = 0;
virtual void train(const arma::mat& batch, Rbm::IListener* pListener) = 0;
arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch);
virtual arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch) = 0;
size_t numTraining();
void addTraining(const arma::mat &toAdd);
+16
View File
@@ -24,6 +24,22 @@ DeepStack::~DeepStack()
{
}
arma::mat DeepStack::trainingBatchFrom(size_t layerId, const arma::mat& batch)
{
arma::mat thisBatch = batch;
Layer *pLayer = getLayer(0);
while (pLayer)
{
if (pLayer->id() == layerId)
{
break;
}
thisBatch = pLayer->toHiddenProbs(thisBatch);
pLayer = pLayer->next;
}
return thisBatch;
}
void DeepStack::train(const arma::mat& batch, Rbm::IListener* pListener)
{
arma::mat thisBatch = batch;
+1
View File
@@ -28,6 +28,7 @@ public:
DeepStack(const DeepStack& orig) = delete;
virtual ~DeepStack();
arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch);
void train(const arma::mat& batch, Rbm::IListener* pListener) override;
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override;
-17
View File
@@ -118,23 +118,6 @@ arma::mat Layer::vc_to_v(const arma::mat& vc) const
return arma::reshape(vc, 1, numVisible() - m_numContext);
}
void Layer::calcContextBatch(arma::mat& batch)
{
size_t numTraining = batch.n_rows;
if (m_numContext > 0 and numTraining > 0)
{
for (int i = 1; i < numTraining; i++)
{
arma::mat v = batch.row(i - 1);
arma::mat h = arma::zeros(1, numHidden());
gibbs_vh(v, h);
arma::mat training_with_ctx = arma::join_rows(batch.row(i).cols(0, numVisible() - m_numContext - 1), h);
batch.row(i) = training_with_ctx;
}
}
}
void Layer::train(const arma::mat& batch, IListener* pListener)
{
std::cout << m_name << ": " << " Training of layer " << std::to_string(id()) << std::endl;
-1
View File
@@ -48,7 +48,6 @@ public:
bool weightsLoad(std::string const &dir, std::string const &prj);
bool weightsSave(std::string const &dir, std::string const &prj);
void calcContextBatch(arma::mat &batch);
void train(arma::mat const &batch, IListener *pListener=nullptr);
arma::mat to_h_gibbs(const arma::mat& v_probs);
+1 -2
View File
@@ -903,8 +903,7 @@ const juce::String& MainComponent::getBaseDir()
void MainComponent::run()
{
trainButton->setButtonText (TRANS("Stop"));
m_pLayer->calcContextBatch(m_stack->trainingBatch());
// m_pLayer->train(m_stack->trainingBatch(), this);
// m_pLayer->calcContextBatch(m_stack->trainingBatch());
m_stack->train(m_pLayer->id(), m_stack->trainingBatch(), this);
trainButton->setButtonText (TRANS("Train"));
}
+73 -3
View File
@@ -13,8 +13,9 @@
#include "RnnStack.hpp"
RnnStack::RnnStack(const std::string &name)
RnnStack::RnnStack(const std::string &name, size_t numContext)
: AStack(StackType::Rnn, name)
, m_numContext(numContext)
{
}
@@ -22,13 +23,82 @@ RnnStack::~RnnStack()
{
}
size_t RnnStack::getSeqLen()
{
return numLayers();
}
size_t RnnStack::numContext()
{
return m_numContext;
}
arma::mat RnnStack::v_to_vc(const arma::mat& v) const
{
arma::mat c = arma::zeros(v.n_rows, m_numContext);
return arma::join_rows(v, c);
}
arma::mat RnnStack::v_to_vc(const arma::mat& v, const arma::mat& c) const
{
return arma::join_rows(v, c);
}
arma::mat RnnStack::vc_to_c(const arma::mat& vc) const
{
size_t numVisible = vc.n_cols;
if (m_numContext == 0)
{
return arma::mat(1, 0);
}
return vc.submat(0, numVisible - m_numContext, 0, numVisible - 1);
}
arma::mat RnnStack::vc_to_v(const arma::mat& vc) const
{
size_t numVisible = vc.n_cols;
return vc.submat(0, 0, vc.n_rows - 1, numVisible - m_numContext - 1);
}
arma::mat RnnStack::trainingBatchFrom(size_t layerId, const arma::mat& batch)
{
arma::mat vc = batch;
Layer *pLayer = getLayer(0);
while (pLayer)
{
if (layerId == pLayer->id())
{
break;
}
arma::mat c = pLayer->to_h_gibbs(vc);
arma::mat v = arma::shift(vc_to_v(vc), layerId, 1);
vc = v_to_vc(v, c);
pLayer = pLayer->next;
}
return vc;
}
void RnnStack::train(const arma::mat& batch, Rbm::IListener* pListener)
{
arma::mat thisBatch = batch;
Layer *pLayer = getLayer(0);
while (pLayer)
{
thisBatch = trainingBatchFrom(pLayer->id(), thisBatch);
pLayer->train(thisBatch, pListener);
pLayer = pLayer->next;
}
}
void RnnStack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener)
{
arma::mat thisBatch = batch;
Layer *pLayer = getLayer(layerId);
if (pLayer)
{
thisBatch = trainingBatchFrom(pLayer->id(), batch);
pLayer->train(thisBatch, pListener);
}
}
+11 -3
View File
@@ -19,16 +19,24 @@
class RnnStack : public AStack
{
public:
RnnStack(const std::string &name);
RnnStack(const std::string &name, size_t numContext);
RnnStack(const RnnStack& orig) = delete;
virtual ~RnnStack();
arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch) override;
void train(const arma::mat& batch, Rbm::IListener* pListener) override;
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override;
size_t getSeqLen();
size_t numContext();
arma::mat v_to_vc(const arma::mat &v) const;
arma::mat v_to_vc(const arma::mat &v, const arma::mat &c) const;
arma::mat vc_to_v(const arma::mat &vc) const;
arma::mat vc_to_c(const arma::mat &vc) const;
private:
size_t m_numContext;
};
#endif /* RNNSTACK_HPP */
+4 -2
View File
@@ -34,11 +34,12 @@ AStack* StackCreator::fromJson(Json::Value& project, LayerConstructor *pLayerCon
Json::Value &stack = project["stack"];
Json::Value &layers = project["stack"]["layers"];
int type = stack["type"].asInt();
int numContext = stack["numContext"].asInt();
AStack *pStack = nullptr;
if (type == AStack::StackType::Rnn)
{
pStack = new RnnStack(name);
pStack = new RnnStack(name, numContext);
}
else
{
@@ -94,6 +95,7 @@ bool StackCreator::toFile(AStack* pStack, const std::string& dir, const std::str
project["stack"]["name"] = name;
project["stack"]["type_string"] = AStack::stackTypeStrings[pStack->type()];
project["stack"]["type"] = pStack->type();
project["stack"]["numContext"] = pStack->numContext();
Json::Value layers(Json::arrayValue);
Layer *pLayer = pStack->getLayer(0);
+19 -5
View File
@@ -9,7 +9,7 @@
#include <jsoncpp/json/json.h>
#include "Rbm.hpp"
#include "Layer.hpp"
#include "DeepStack.hpp"
#include "RnnStack.hpp"
#include "StackCreator.hpp"
using namespace std;
@@ -140,11 +140,11 @@ arma::mat createTraining(const string &filename, size_t seq_len)
struct Rnn
{
Rnn(Layer *layer)
: nV(layer->numVisible() - layer->context().n_cols)
: nV(layer->numVisible() - layer->numHidden())
, nH(layer->numHidden())
, nVx(layer->numVisibleX())
, nVy(layer->numVisibleY())
, nC(layer->context().n_cols)
, nC(layer->numHidden())
{
}
@@ -192,6 +192,7 @@ arma::mat to_next(arma::mat v)
}
#define CREATE_TRAINING 0
#define DO_TRAINING 1
int main()
{
#if CREATE_TRAINING
@@ -201,7 +202,7 @@ int main()
#endif
// Load project
AStack *stack = StackCreator::fromFile(".", "poet5");
RnnStack *stack = reinterpret_cast<RnnStack*>(StackCreator::fromFile(".", "poet2"));
// Load weights
stack->loadWeights(".");
@@ -209,6 +210,19 @@ int main()
// Load training
stack->loadTrainingBatch(".");
#if DO_TRAINING
RbmListener listener;
arma::mat t_vc = stack->trainingBatch();
for (int i=0; i < stack->getSeqLen(); i++)
{
stack->getLayer(i)->weightsInit(0.1,0);
stack->train(i, t_vc, &listener);
}
stack->saveWeights(".");
stack->save(".");
#else
Layer *layer = stack->getLayer(0);
int numTraining = stack->trainingBatch().n_rows;
@@ -268,7 +282,7 @@ int main()
putchar(c);
}
#endif
printf("\n\nEnd of program\n");
return 0;
}