- integrated RnnStack
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@826 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+5
-1
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+9
-1
@@ -19,15 +19,23 @@
|
||||
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;
|
||||
|
||||
};
|
||||
|
||||
|
||||
@@ -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
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user