- refactored
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@820 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
@@ -39,6 +39,11 @@ AStack::~AStack()
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
AStack::StackType AStack::type()
|
||||||
|
{
|
||||||
|
return m_type;
|
||||||
|
}
|
||||||
|
|
||||||
void AStack::setName(const std::string& name)
|
void AStack::setName(const std::string& name)
|
||||||
{
|
{
|
||||||
m_name = name;
|
m_name = name;
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ public:
|
|||||||
AStack(const AStack& orig);
|
AStack(const AStack& orig);
|
||||||
virtual ~AStack();
|
virtual ~AStack();
|
||||||
|
|
||||||
|
StackType type();
|
||||||
void setName(const std::string &name);
|
void setName(const std::string &name);
|
||||||
size_t numLayers();
|
size_t numLayers();
|
||||||
void addLayer(Layer *pLayer);
|
void addLayer(Layer *pLayer);
|
||||||
|
|||||||
@@ -29,7 +29,13 @@ public:
|
|||||||
virtual ~DeepStack();
|
virtual ~DeepStack();
|
||||||
|
|
||||||
void train(Rbm::IListener* pListener);
|
void train(Rbm::IListener* pListener);
|
||||||
|
void train(const arma::mat& batch, Rbm::IListener* pListener) override;
|
||||||
|
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override;
|
||||||
|
|
||||||
|
arma::mat upPass(size_t layerId, arma::mat const &v);
|
||||||
|
arma::mat downPass(size_t layerId, arma::mat const &h);
|
||||||
|
arma::mat upDownPass(size_t layerId, arma::mat const &v);
|
||||||
|
|
||||||
size_t numTraining();
|
size_t numTraining();
|
||||||
void addTraining(const arma::mat &toAdd);
|
void addTraining(const arma::mat &toAdd);
|
||||||
void delTraining(int index);
|
void delTraining(int index);
|
||||||
@@ -38,13 +44,6 @@ public:
|
|||||||
arma::mat& trainingBatch();
|
arma::mat& trainingBatch();
|
||||||
arma::mat trainingBatch(Layer *pLayer);
|
arma::mat trainingBatch(Layer *pLayer);
|
||||||
|
|
||||||
arma::mat upPass(size_t layerId, arma::mat const &v);
|
|
||||||
arma::mat downPass(size_t layerId, arma::mat const &h);
|
|
||||||
arma::mat upDownPass(size_t layerId, arma::mat const &v);
|
|
||||||
|
|
||||||
void train(const arma::mat& batch, Rbm::IListener* pListener) override;
|
|
||||||
void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override;
|
|
||||||
|
|
||||||
private:
|
private:
|
||||||
arma::mat m_trainingBatch;
|
arma::mat m_trainingBatch;
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -49,7 +49,6 @@ public:
|
|||||||
bool weightsSave(std::string const &dir, std::string const &prj);
|
bool weightsSave(std::string const &dir, std::string const &prj);
|
||||||
|
|
||||||
void calcContextBatch(arma::mat &batch);
|
void calcContextBatch(arma::mat &batch);
|
||||||
arma::mat trainingData(arma::mat const &batch);
|
|
||||||
void train(arma::mat const &batch, IListener *pListener=nullptr);
|
void train(arma::mat const &batch, IListener *pListener=nullptr);
|
||||||
|
|
||||||
arma::mat to_h_gibbs(const arma::mat& v_probs);
|
arma::mat to_h_gibbs(const arma::mat& v_probs);
|
||||||
|
|||||||
Reference in New Issue
Block a user