- refactored

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@820 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-17 15:14:55 +00:00
parent dbed14cefa
commit 370fdd7bbb
4 changed files with 12 additions and 8 deletions
+5
View File
@@ -39,6 +39,11 @@ AStack::~AStack()
}
}
AStack::StackType AStack::type()
{
return m_type;
}
void AStack::setName(const std::string& name)
{
m_name = name;
+1
View File
@@ -50,6 +50,7 @@ public:
AStack(const AStack& orig);
virtual ~AStack();
StackType type();
void setName(const std::string &name);
size_t numLayers();
void addLayer(Layer *pLayer);
+6 -7
View File
@@ -29,7 +29,13 @@ public:
virtual ~DeepStack();
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();
void addTraining(const arma::mat &toAdd);
void delTraining(int index);
@@ -38,13 +44,6 @@ public:
arma::mat& trainingBatch();
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:
arma::mat m_trainingBatch;
};
-1
View File
@@ -49,7 +49,6 @@ public:
bool weightsSave(std::string const &dir, std::string const &prj);
void calcContextBatch(arma::mat &batch);
arma::mat trainingData(arma::mat const &batch);
void train(arma::mat const &batch, IListener *pListener=nullptr);
arma::mat to_h_gibbs(const arma::mat& v_probs);